Particle Life
¶
Installation¶
You will need Python 3.12 or later, and a working JAX installation. For example, you can install JAX with:
In [1]:
Copied!
%pip install -U "jax[cuda]"
%pip install -U "jax[cuda]"
/home/faldor_google_com/dev/cax/.venv/bin/python3: No module named pip
Note: you may need to restart the kernel to use updated packages.
Then, install CAX from PyPi:
In [2]:
Copied!
%pip install -U "cax[examples]"
%pip install -U "cax[examples]"
/home/faldor_google_com/dev/cax/.venv/bin/python3: No module named pip
Note: you may need to restart the kernel to use updated packages.
Import¶
In [3]:
Copied!
import jax
import jax.numpy as jnp
import mediapy
from flax import nnx
from cax.cs.particle_life import ParticleLife, ParticleLifeState
from cax.utils import render_states
import jax
import jax.numpy as jnp
import mediapy
from flax import nnx
from cax.cs.particle_life import ParticleLife, ParticleLifeState
from cax.utils import render_states
Configuration¶
In [4]:
Copied!
seed = 0
num_steps = 1024
num_spatial_dims = 2
num_particles = 4096
num_classes = 6
dt = 0.01
force_factor = 1.0
velocity_half_life = dt
r_max = 0.15
beta = 0.3
key = jax.random.key(seed)
rngs = nnx.Rngs(seed)
seed = 0
num_steps = 1024
num_spatial_dims = 2
num_particles = 4096
num_classes = 6
dt = 0.01
force_factor = 1.0
velocity_half_life = dt
r_max = 0.15
beta = 0.3
key = jax.random.key(seed)
rngs = nnx.Rngs(seed)
Instantiate system¶
In [5]:
Copied!
# Sample attraction matrix
key, subkey = jax.random.split(key)
attraction_matrix = jax.random.uniform(
subkey, (num_classes, num_classes), minval=-1.0, maxval=1.0
)
attraction_matrix
# Sample attraction matrix
key, subkey = jax.random.split(key)
attraction_matrix = jax.random.uniform(
subkey, (num_classes, num_classes), minval=-1.0, maxval=1.0
)
attraction_matrix
Out[5]:
Array([[-0.98541236, -0.9582176 , 0.162853 , -0.27632403, -0.55392456,
-0.76142335],
[-0.74912214, 0.23365998, -0.8099854 , 0.96684504, -0.15035582,
0.66496015],
[-0.8067758 , -0.5459583 , -0.2908504 , 0.34208274, -0.7660713 ,
-0.05521393],
[-0.6609576 , -0.9498596 , -0.731271 , -0.25255132, 0.8165903 ,
0.5833442 ],
[-0.7534721 , -0.41303635, 0.58081174, 0.9019952 , -0.31768775,
0.0429759 ],
[-0.48326612, -0.9228716 , -0.16752648, 0.64369607, -0.8060143 ,
-0.5253153 ]], dtype=float32)
In [6]:
Copied!
cs = ParticleLife(
num_classes=num_classes,
dt=dt,
force_factor=force_factor,
velocity_half_life=velocity_half_life,
r_max=r_max,
beta=beta,
attraction_matrix=attraction_matrix,
)
cs = ParticleLife(
num_classes=num_classes,
dt=dt,
force_factor=force_factor,
velocity_half_life=velocity_half_life,
r_max=r_max,
beta=beta,
attraction_matrix=attraction_matrix,
)
Sample initial state¶
In [7]:
Copied!
def sample_state(key):
"""Sample a state with random classes and positions, and zero velocity."""
key_class, key_position = jax.random.split(key)
# Class
class_id = jax.random.choice(key_class, num_classes, (num_particles,))
# Position
position = jax.random.uniform(
key_position, (num_particles, num_spatial_dims), minval=0.0, maxval=1.0
)
# Velocity
velocity = jnp.zeros((num_particles, num_spatial_dims))
return ParticleLifeState(class_id=class_id, position=position, velocity=velocity)
def sample_state(key):
"""Sample a state with random classes and positions, and zero velocity."""
key_class, key_position = jax.random.split(key)
# Class
class_id = jax.random.choice(key_class, num_classes, (num_particles,))
# Position
position = jax.random.uniform(
key_position, (num_particles, num_spatial_dims), minval=0.0, maxval=1.0
)
# Velocity
velocity = jnp.zeros((num_particles, num_spatial_dims))
return ParticleLifeState(class_id=class_id, position=position, velocity=velocity)
Run¶
In [8]:
Copied!
key, subkey = jax.random.split(key)
state_init = sample_state(subkey)
state_final, states = cs(state_init, num_steps=num_steps, return_states=True)
key, subkey = jax.random.split(key)
state_init = sample_state(subkey)
state_final, states = cs(state_init, num_steps=num_steps, return_states=True)
Visualize¶
In [9]:
Copied!
states = jax.tree.map(lambda x, xs: jnp.concatenate([x[None], xs]), state_init, states)
frames = render_states(cs, states, particle_radius=0.003)
mediapy.show_video(frames, fps=int(1 / dt))
states = jax.tree.map(lambda x, xs: jnp.concatenate([x[None], xs]), state_init, states)
frames = render_states(cs, states, particle_radius=0.003)
mediapy.show_video(frames, fps=int(1 / dt))