From f17b181781152c508fa8fa35c056b5be54b49493 Mon Sep 17 00:00:00 2001 From: Will Dean Date: Mon, 5 Oct 2026 12:10:06 -0400 Subject: [PATCH] Fix dims type in model_from_fgraph, and test dims preservation add_named_variable is declared to take a tuple and was passed a list. It stores dims verbatim, so the list reached named_vars_to_dims and the round trip stopped agreeing with itself: >>> model_from_fgraph(fgraph_from_model(m)[0]).named_vars_to_dims == m.named_vars_to_dims False # before True # after Every path built on the pair inherits it, so Model.copy(), copy.copy, copy.deepcopy, pm.do and pm.observe all disagreed with their input too. The tests were part of the defect. They asserted the round trip's own output, ["test_dim"], rather than comparing to the original model, which is why this went unnoticed. They now compare against the original, and that is the assertion that catches the type. Invisible to CI because both files sit in the FAILING mypy baseline. No consumer depends on a list here: backends/arviz.py and backends/mcbackend.py normalise with list(), printing.py and variational/opvi.py ignore the type, and testing.py never reads dims. So this is a contract fix, not a crash fix. --- pymc/model/fgraph.py | 4 ++-- tests/model/test_fgraph.py | 9 +++++---- tests/model/transform/test_conditioning.py | 9 ++++++--- 3 files changed, 13 insertions(+), 9 deletions(-) diff --git a/pymc/model/fgraph.py b/pymc/model/fgraph.py index 404df717c2..805208d7e0 100644 --- a/pymc/model/fgraph.py +++ b/pymc/model/fgraph.py @@ -375,10 +375,10 @@ def first_non_model_var(var): [var] = model_var.owner.inputs model.data_vars.append(var) else: - raise TypeError(f"Unexpected ModelVar type {type(model_var)}") + raise TypeError(f"Unexpected ModelVar type {type(op)}") var.name = op.name - dims = list(op.dims) if op.dims else None + dims = tuple(op.dims) if op.dims else None model.add_named_variable(var, dims=dims) return model diff --git a/tests/model/test_fgraph.py b/tests/model/test_fgraph.py index 13a1a3e9de..86d9274e2d 100644 --- a/tests/model/test_fgraph.py +++ b/tests/model/test_fgraph.py @@ -61,7 +61,8 @@ def test_basic(): assert m_new.coords == {"test_dim": tuple(range(3))} assert m_new._dim_lengths["test_dim"].eval() == 3 - assert m_new.named_vars_to_dims == {"z": ["test_dim"]} + assert m_new.named_vars_to_dims == {"z": ("test_dim",)} + assert m_new.named_vars_to_dims == m_old.named_vars_to_dims named_vars = {"x", "y", "w", "z", "pot"} assert set(m_new.named_vars) == named_vars @@ -353,9 +354,9 @@ def test_fgraph_rewrite(non_centered_rewrite): m_new = model_from_fgraph(fg) assert m_new.named_vars_to_dims == { - "subject_mean": ["subject"], - "subject_mean_raw_": ["subject"], - "obs": ["subject"], + "subject_mean": ("subject",), + "subject_mean_raw_": ("subject",), + "obs": ("subject",), } assert set(m_new.named_vars) == { "group_mean", diff --git a/tests/model/transform/test_conditioning.py b/tests/model/transform/test_conditioning.py index 5f8eb99ee4..d4c2235135 100644 --- a/tests/model/transform/test_conditioning.py +++ b/tests/model/transform/test_conditioning.py @@ -132,7 +132,8 @@ def test_observe_dims(): x = pm.Normal("x", dims="test_dim") m_new = observe(m_old, {x: np.arange(5, dtype=config.floatX)}) - assert m_new.named_vars_to_dims["x"] == ["test_dim"] + assert m_new.named_vars_to_dims["x"] == ("test_dim",) + assert m_new.named_vars_to_dims == m_old.named_vars_to_dims def test_do(): @@ -229,13 +230,15 @@ def test_do_dims(): m, {"x": np.zeros(10, dtype=config.floatX)}, ) - assert do_m.named_vars_to_dims["x"] == ["test_dim"] + assert do_m.named_vars_to_dims["x"] == ("test_dim",) + assert do_m.named_vars_to_dims == m.named_vars_to_dims do_m = do( m, {"y": np.zeros(10, dtype=config.floatX)}, ) - assert do_m.named_vars_to_dims["y"] == ["test_dim"] + assert do_m.named_vars_to_dims["y"] == ("test_dim",) + assert do_m.named_vars_to_dims == m.named_vars_to_dims @pytest.mark.parametrize("prune", (False, True))