Boids
cax.cs.boids.cs.Boids
Bases: ComplexSystem[BoidsState, Array]
Boids class.
Source code in src/cax/cs/boids/cs.py
24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 | |
__init__(*, dt=0.01, velocity_half_life=jnp.inf, boid_policy)
Initialize Boids.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
dt
|
float
|
Time step of the simulation in arbitrary time units. Smaller values produce smoother motion but require more steps for the same duration. |
0.01
|
velocity_half_life
|
float
|
Time constant for velocity decay due to friction. After this time, velocity is halved without steering input. Use jnp.inf for no friction. Smaller values create more damped, sluggish motion. |
inf
|
boid_policy
|
BoidsPolicy
|
Policy defining the behavior of the boids. |
required |
Source code in src/cax/cs/boids/cs.py
27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 | |
render(state, *, resolution=512, particle_radius=0.01)
Render state to RGB image.
Renders boids as triangular agents pointing in their direction of motion on a white background. Each boid is drawn as a filled triangle with the tip pointing in the velocity direction, providing visual feedback about both position and heading. The visualization uses 2D coordinates in the range [0, 1].
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
state
|
BoidsState
|
BoidsState containing position and velocity arrays. Position should have shape (num_boids, 2) with coordinates in [0, 1]. Velocity determines the orientation and can have arbitrary magnitude. |
required |
resolution
|
int
|
Size of the output image in pixels for both width and height. Higher values produce smoother, more detailed renderings. |
512
|
particle_radius
|
float
|
Half-extent of each boid glyph in coordinate space [0, 1]: the triangle spans twice this value from base to tip and one radius across the base. Larger values make boids more visible but may cause overlap. |
0.01
|
Returns:
| Type | Description |
|---|---|
Array
|
RGB image with dtype uint8 and shape (resolution, resolution, 3), where boids appear as black triangles on a white background. |
Source code in src/cax/cs/boids/cs.py
60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 | |
__call__(state, input=None, *, num_steps=1, input_in_axis=None, return_states=False)
Step the system for multiple time steps.
This method wraps _step inside a JAX scan for efficiency and JIT-compiles the
loop. If input is time-varying, set input_in_axis to the axis containing the
time dimension so that each step receives the corresponding slice of input.
Under return_states=True, the per-step states are also returned as the scan's
stacked outputs, mirroring the (carry, ys) convention of jax.lax.scan. The
trajectory holds the state after each step, stacked along a new leading axis
of size num_steps — its first element is the state after one step, its last
equals the final state, and the initial state is not included.
When remat is enabled, the scan body is wrapped with nnx.remat to reduce
memory usage during backpropagation at the cost of recomputing intermediates.
Note that num_steps, input_in_axis, and return_states are static: each
distinct combination compiles once, so sweeps over horizons should batch their
step counts.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
state
|
State
|
Current state. |
required |
input
|
Input | None
|
Optional input. |
None
|
num_steps
|
int
|
Number of steps. |
1
|
input_in_axis
|
int | None
|
Axis for input if provided for each step. |
None
|
return_states
|
bool
|
Whether to also return the stacked per-step states. |
False
|
Returns:
| Type | Description |
|---|---|
State | tuple[State, State]
|
Final state after |
Source code in src/cax/core/cs.py
57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 | |