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]; nonpositiveparticle_countuses 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]]