NASimJax: A GPU-Accelerated Policy Learning Framework for Penetration Testing

arXiv:2603.19864 · cs.LG, cs.CR · Submitted 2026-03-20 · Read on arXiv

Listen

Radio episode about this paper

Transcript

Introduction to the show: ident: AI Radio. Generated commentary on the latest Artificial Intelligence papers.

Tom: I'm Tom, and with me are Jane, Lu, senior AI researcher at Tsinghua, Meng, lead engineer at a mysterious AI startup and Lalam, the in-house Large Language Model.

Jane: Today's paper: "NASimJax: A GPU-Accelerated Policy Learning Framework for Penetration Testing".

Tom: NASimJax introduces a JAX-based framework that accelerates policy learning for penetration testing by achieving up to 100× higher environment throughput than original simulators.

Jane: First, who's behind it and why it matters.

Title and authors: Tom: So, looking at the authors and that title again, "NASimJax: A GPU-Accelerated Policy Learning Framework for Penetration Testing," it’s clear they are putting their focus squarely on solving a very specific problem in AI research—the gap between slow simulation and scalable policy training. It sounds like they are presenting a concrete tool to bridge that gap.

Jane: Exactly, Tom; the authors clearly recognized that existing simulators were creating a major hurdle for applying reinforcement learning to penetration testing because they couldn't keep up with the millions of interactions needed for effective training on realistic networks. This paper introduces NASimJax as their solution, focusing on making it run much faster using specialized hardware.

Lu: I find the combination of JAX-based RL implementation and the Contextual POMDP formulation really intriguing because it suggests a unified approach: improving simulation speed while simultaneously providing a theoretical structure that helps policies generalize better across different network contexts.

Meng: From a practical standpoint, achieving up to one hundred times higher throughput is substantial; it means we could potentially test vastly larger and more diverse network structures within the same computational resources we currently have available <ref:2603.19864#pg0>. That capability directly impacts how much complexity we can afford to explore in our models.

Lalam: I think the core implication here for AI culture is showing that theoretical frameworks, like contextual POMDPs, combined with efficient hardware utilization, can lead to practical tools that allow researchers to push the boundaries of what's possible in complex domains like cybersecurity simulation.

The paper's summary: Tom: Now we move into the summary of NASimJax: they’re essentially saying they took the old Network Attack Simulator, renamed it NASimJax, and completely rebuilt it using JAX on accelerators to achieve that significant speedup. They also introduced a new network generation pipeline designed to create scenarios that are both structurally diverse and guaranteed to be solvable.

Jane: That part about the network generation pipeline sounds really important because if the environments they generate aren't diverse enough, any policy trained on them won't work well when faced with a completely new kind of attack scenario in the real world. They’ve built a system that aims for variety while keeping things manageable.

Lu: The formulation as a Contextual POMDP means that each episode is conditioned on context describing the network instance, which is a powerful concept for studying how an agent should adapt its strategy based on what it observes about the underlying network structure before it even starts acting.

Meng: I'm looking at that guaranteed solvability aspect; if they can ensure every sensitive host has a path to privilege escalation, that simplifies the training problem immensely because we don't have to worry about environments being impossible to solve by accident. That’s a huge practical win for training stability.

Lalam: The summary highlights how this approach provides a principled basis for studying zero-shot policy generalization, which is precisely what we need when agents are deployed in environments they haven't seen before. It moves the discussion from just making simulations run faster to creating better learning theory.

The paper's improvements: Tom: The paper details several key improvements they made, and I want to talk about the two main ways they handled the growing action space in larger networks: Action Masking and Two-Stage Action Selection, or 2SAS <ref:2603.19864#pg0>. They argue that 2SAS is actually better than just flat masking when dealing with very large networks <ref:2603.19864#pg0>.

Jane: That distinction between simple masking and the decomposition offered by 2SAS is something I think we should unpack because it shows they didn't just apply a quick fix to the action space; they built a more sophisticated mechanism for managing complexity <ref:2603.19864#pg0>. It breaks down the decision-making process into smaller, manageable steps.

Lu: The idea behind 2SAS, inspired by factored action decomposition, is clever because it separates the choice of which host to attack from what specific exploit or privilege escalation step to take on that host <ref:2603.19864#pg0>. This separation should make the policy learning much more efficient when the number of hosts gets large.

Meng: As an engineer, I'm focused on how this translates into actual computational efficiency; if 2SAS reduces the effective decision complexity substantially, that means we can afford to train policies for networks with more than just a few hosts, which is where the real scaling happens <ref:2603.19864#pg0>.

Lalam: The authors also introduced several training techniques like Prioritized Level Replay and Domain Randomization, alongside reward scaling to handle variance. These aren't just tweaks; they form a cohesive strategy designed to build a curriculum that helps the agent learn robust policies across different network contexts.

Conclusion: Tom: So, wrapping up this discussion on "NASimJax: A GPU-Accelerated Policy Learning Framework for Penetration Testing," the main points are that they achieved a one hundred times faster environment throughput using JAX acceleration and formalized the problem as a Contextual POMDP <ref:2603.19864#pg0,NASimJax: A GPU-Accelerated Policy Learning Framework for Penetration Testing>. They also showed that advanced techniques like 2SAS handle large action spaces better than simpler masking, and they used curriculum learning strategies to improve generalization <ref:2603.19864#pg0>.

Jane: It really shows how structuring the RL problem correctly, by defining it as a contextual POMDP, combined with smart architectural choices like 2SAS, can lead to policies that are much more robust when applied to unseen network architectures <ref:2603.19864#pg0>. The implications for testing security systems in the real world are substantial because it means we can train agents on environments that mirror real-world complexity at a speed previously thought impossible.

Lu: I think the potential here is immense because if we can reliably create these structurally diverse and guaranteed-solvable scenarios, we start having a principled way to study how policies transfer from training data to completely novel network topologies, which is a core challenge in AI generalization.

Meng: From an engineering view, the ability to scale policy training up by that factor means we can move from testing small proof-of-concept networks to exploring much more realistic, large-scale infrastructure models under fixed compute constraints. That’s where the practical application lies for us.

Lalam: Ultimately, this work contributes a framework that allows AI agents to learn policies that are not just good at one specific network layout but are actually capable of generalizing their attack strategies across a distribution of possible environments. It elevates the entire field by providing a solid foundation for scalable, context-aware learning in security tasks.

CISS Department, Royal Military Academy, Belgium · AI Lab, Vrije Universiteit Brussel, Belgium

cs.LG, cs.CR

Submitted: 2026-03-20

Updated: 2026-10-04

Code: https://github.com/raphsimon/NASimJax

Importance score: 89/100

The gist: NASimJax introduces a JAX-based framework that accelerates policy learning for penetration testing by achieving up to 100× higher environment throughput than original simulators.

Key concepts

Contextual POMDP
This models automated penetration testing where the current situation (context) describes the underlying network instance. The agent learns a policy based on this context, allowing it to generalize its actions across many different network setups, which is key for zero-shot learning.
Network Generation Pipeline
This process creates realistic and diverse test environments by carefully controlling host properties (OS, services), sensitivity, and connectivity. Constraints ensure that every sensitive machine has exploitable vulnerabilities, creating a curriculum of varying difficulty for the learning agent.
Action Masking
This technique reduces complexity in large networks by virtually eliminating invalid actions. It masks choices for undiscovered hosts or exploits that require missing software/processes, ensuring the agent only considers feasible and relevant moves during decision-making.
Two-Stage Action Selection (2SAS)
Instead of one complex decision, 2SAS splits the action into two parts: first selecting a host, then selecting an exploit for that specific host. This decomposition manages large action spaces by making decisions sequentially, significantly reducing the effective complexity at each step.

Terminology

Summary

NASimJax introduces a JAX-based framework that accelerates policy learning for penetration testing by achieving up to 100× higher environment throughput than original simulators. This framework reformulates automated penetration testing as a Contextual POMDP and introduces novel network generation pipelines to produce structurally diverse and guaranteed-solvable scenarios, providing a principled basis for studying zero-shot policy generalization.

The gist

NASimJax achieves up to 100× higher environment throughput than the original simulator by running the entire training pipeline on hardware accelerators, enabling experimentation on larger networks under fixed compute budgets that were previously infeasible.

How it works

NASimJax models automated penetration testing as a Contextual POMDP, where each episode is conditioned on a context describing the underlying network instance. This formulation provides a principled framework for training policies across a distribution of environments and directly facilitates zero-shot policy generalization to previously unseen networks. The environment adheres to the Gymnax API, enabling seamless integration with JAX-based RL algorithms.

Network Generation Pipeline

The network generation process is designed to balance realism, diversity, and guaranteed solvability through several stages:

  1. Initially, a fixed number of hosts are distributed across subnets with two special semantic roles: the Internet subnet containing the attacker’s machine and a demilitarized zone (DMZ).

  2. Host-level properties are assigned randomly according to fixed distributions for operating systems, service density (svcd), and process density (procd). Host sensitivity is assigned independently with probability sd.

  3. Feasibility constraints are enforced: every subnet must contain at least one host running at least one service, and every sensitive host must run at least one service and process to ensure privilege escalation is possible on all sensitive machines.

  4. Network topology is then generated via an adjacency matrix of size Ns × Ns, where entries are sampled according to a topology density parameter td. Connectivity between subnets is directed, with the Internet subnet restricted to communicating only with the DMZ. This asymmetry can lead to subnets that become inactive for a given episode, shaping the learning problem and inducing a curriculum of environments with varying difficulty levels.

Handling Large Action Spaces

To address the linearly growing action space of larger networks, NASimJax introduces two distinct methods:

  1. Action Masking: This method modifies the categorical distribution by giving invalid actions an infinitesimal value before applying softmax, resulting in actions whose sampling probability becomes virtually zero. It masks all actions for hosts that have both not been discovered yet and are not reachable, and also masks exploits and privilege escalation actions that are invalid due to missing OS/service or OS/process combinations.

  2. Two-Stage Action Selection (2SAS): Inspired by factored action decomposition, this mechanism decomposes each decision into two stages: host selection followed by per-host action selection. The actor-critic network branches into two policy heads: the first outputs a distribution over hosts, masked to exclude unreachable or undiscovered targets; the second takes a learned host embedding and outputs a distribution over per-host actions, masked to remove invalid choices. This reduces effective decision complexity at each stage substantially.

Training and Generalization Methods

The framework employs several techniques to improve generalization across network contexts:

  1. Prioritized Level Replay (PLR): PLR works in conjunction with procedurally generated environments, aiming to form a natural curriculum of levels where the agent exhibits the highest regret with respect to its current policy. It performs gradient updates on random levels, whereas robust PLR (PLR⊥) only updates the agent on levels sampled from the buffer.

  2. Domain Randomization (DR): DR involves generating a new network from a parameterized distribution at each end of an episode.

  3. Reward Scaling: To mitigate high variance in cumulative reward due to varying numbers of sensitive hosts, rewards are scaled such that the maximum potential return is approximately invariant to the network size: rˆt = rt / (Ns·Vh), where Vh is the reward value for compromising a sensitive host. This ensures that learning signals reflect structural difficulty rather than network size alone.

Policy Learning Algorithms

The implementation is based on a pure JAX PPO implementation from Lu et al. [2022]. To handle partial observability, the methodology follows Simon et al. [2025], where the most recent observation is kept and the accumulated history of past observations is appended to form the state representation. The training process utilizes hyperparameter tuning via Bayesian-sampling for every algorithm, with a budget of 250 trials, allowing for robust selection of parameters across different network sizes and algorithms. The experiments investigate three primary questions: speed-up achieved by NASimJax, which reaches 1.6M steps per second on a single entry-level GPU; the effectiveness of action space handling methods (showing 2SAS outperforms flat masking on larger networks); and the performance of UED methods on zero-shot policy transfer (ZSPT)

Improvements for AI systems

As a fastidious and diligent researcher, my analysis of NASimJax reveals several high-impact areas for improvement in AI systems, specifically in the domain of Reinforcement Learning (RL) applied to complex cyber-physical or sequential decision-making tasks like penetration testing.

Here are the specific improvements and what the enhanced AI system can achieve:


The core improvement lies in transitioning from slow, fixed simulation environments to a high-throughput, flexible Contextual POMDP framework that enables robust generalization across diverse network topologies.

  1. A complete overhaul of the RL training infrastructure from CPU-bound Python implementations to a JAX-based, GPU-accelerated pipeline (NASimJax).

  2. Implementation of novel curriculum learning strategies using Prioritized Level Replay (PLR) combined with Domain Randomization (DR) or contextual network generation, allowing agents to learn effective policies on sparser topologies first, which implicitly builds competence for denser ones.

  3. Integration of a two-stage action decomposition (2SAS) mechanism to handle linearly growing action spaces efficiently, significantly outperforming flat action masking at scale.

  4. A principled reward scaling strategy that normalizes rewards based on the theoretical maximum return per context, ensuring stable training signals across varied network sizes and densities.

The improved AI system (NASimJax-based RL Agent) can achieve the following specific capabilities:

  1. An agent capable of learning generalized penetration testing policies that perform well on unseen network architectures (Zero-Shot Policy Transfer).

  2. The ability to scale policy training to much larger networks and more complex attack scenarios than previously feasible due to a 100x speedup in environment throughput, allowing for comprehensive exploration under fixed compute budgets.

  3. Robust performance across a wide distribution of network densities (topology parameters like connectivity, service density, and host sensitivity), meaning the agent won't overfit to a single typical network structure encountered during training.

  4. Efficient decision-making in environments with massive action spaces (e.g., 40+ hosts) by intelligently decomposing complex actions into host selection and per-host actions (2SAS), leading to superior solve rates compared to simpler masking techniques.

  5. Systematic investigation of the failure modes between replay mechanisms (PLR) and credit assignment structures (2SAS), providing concrete insights into why certain RL architectures fail under specific scaling conditions, leading to more resilient system designs.

Sources

Related papers