def make_rigid_body_composition[State](
state: State,
cloud: PointCloud[MCMCParticles, MCMCSystems],
motifs: Table[MotifParticleId, MotifParticles],
groups: Lens[State, Buffered[GroupId, MCMCGroup]],
) -> RigidBodyComposition[State, MotifParticles]:
"""Describe rigid bodies that can be inserted into a fixed host."""
systems = cloud.systems
accept = systems.set_data(jnp.ones(len(systems), dtype=bool))
def counts(
state: State, patch: Patch[State] | None, old_input: bool = False
) -> Table[SystemId, Array]:
if patch is not None and not old_input:
state = patch(state, accept)
values = motif_counts(groups(state))
return systems.set_data(
jnp.concatenate((jnp.ones((len(systems), 1)), values), axis=1)
)
particles = cloud.particles.data
return RigidBodyComposition(
state,
particles.system.apply_mask(~particles.group.valid_mask),
motifs.data,
motifs.data.motif,
counts,
)