waxmorph.jax.graph.build_node_features

waxmorph.jax.graph.build_node_features#

waxmorph.jax.graph.build_node_features(c, particle_count)[source]#

Return concentrations as float32 [N, C] node features.

Rank-one input becomes [N, 1]; nonpositive particle_count uses every row.

Examples

>>> c = jnp.array([0.2, 0.4, 0.8])
>>> print(build_node_features(c, 3).tolist())
[[0.20000000298023224], [0.4000000059604645], [0.800000011920929]]
Parameters:

particle_count (int)

Return type:

Array