Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions pymc/model/fgraph.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
9 changes: 5 additions & 4 deletions tests/model/test_fgraph.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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",
Expand Down
9 changes: 6 additions & 3 deletions tests/model/transform/test_conditioning.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand Down Expand Up @@ -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))
Expand Down
Loading