From da9b7bfdd8a09f501fe041abc43d8e396f74caec Mon Sep 17 00:00:00 2001 From: zombie-einstein <13398815+zombie-einstein@users.noreply.github.com> Date: Tue, 11 Nov 2025 23:13:09 +0000 Subject: [PATCH] Use tuple dims in generator --- pyproject.toml | 2 +- src/thants/generators/colonies/utils.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 4a74e5e..1c371bc 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "thants" -version = "0.3.0" +version = "0.3.1" description = "Ants! On your GPU! With JAX!" authors = ["zombie-einstein "] license = "MIT" diff --git a/src/thants/generators/colonies/utils.py b/src/thants/generators/colonies/utils.py index c8ef9bf..e1f1931 100644 --- a/src/thants/generators/colonies/utils.py +++ b/src/thants/generators/colonies/utils.py @@ -80,10 +80,10 @@ def init_colony( nest_idxs = get_rectangular_indices(nest_dims) nest_idxs = nest_idxs + centre - jnp.array(nest_dims) // 2 nest_idxs = nest_idxs % dims_arr - nest = jnp.zeros(dims_arr, dtype=bool) + nest = jnp.zeros(dims, dtype=bool) nest = nest.at[nest_idxs[:, 0], nest_idxs[:, 1]].set(True) - signals = jnp.zeros((n_signals, *dims_arr)) + signals = jnp.zeros((n_signals, *dims)) return Colony(ants=ants, signals=signals, nest=nest)