|
22 | 22 | import pytensor.tensor as pt |
23 | 23 | import pytest |
24 | 24 |
|
| 25 | +from xarray import Dataset |
| 26 | + |
25 | 27 | import pymc as pm |
26 | 28 | import pymc.variational.opvi as opvi |
27 | 29 |
|
@@ -509,9 +511,163 @@ def test_sample_outside_model_context(): |
509 | 511 | with pm.Model() as model: |
510 | 512 | mu = pm.Normal("mu", 0, 1) |
511 | 513 |
|
512 | | - # Fit with explicit model, then exit the model context |
513 | 514 | approx = pm.fit(10, method="advi", model=model, progressbar=False) |
514 | | - |
515 | | - # sample() should work without an active model context |
516 | 515 | trace = approx.sample(50) |
517 | 516 | assert trace.posterior["mu"].shape == (1, 50) |
| 517 | + |
| 518 | + |
| 519 | +class TestUntransformedData: |
| 520 | + def test_state_mean_field(self): |
| 521 | + """ADVI state has family='mean_field', mean and std in constrained space.""" |
| 522 | + rng = np.random.default_rng(42) |
| 523 | + with pm.Model(): |
| 524 | + pm.HalfNormal("sigma", sigma=5.0) |
| 525 | + pm.Normal("mu", 0, 1) |
| 526 | + pm.Normal("y", rng.normal(size=3), observed=rng.normal(size=3)) |
| 527 | + fitted = pm.fit(100, method="advi", progressbar=False, random_seed=42) |
| 528 | + |
| 529 | + s = fitted.state |
| 530 | + assert set(s.mean.keys()) == {"sigma", "mu"} |
| 531 | + assert set(s.std.keys()) == {"sigma", "mu"} |
| 532 | + assert s.std is not None |
| 533 | + assert s.mean["sigma"].values > 0 |
| 534 | + assert s.std["sigma"].values > 0 |
| 535 | + |
| 536 | + def test_state_full_rank(self): |
| 537 | + """FullRankADVI state has mean and std.""" |
| 538 | + rng = np.random.default_rng(42) |
| 539 | + with pm.Model(): |
| 540 | + pm.HalfNormal("sigma", sigma=5.0) |
| 541 | + pm.Normal("mu", 0, 1) |
| 542 | + pm.Normal("y", rng.normal(size=3), observed=rng.normal(size=3)) |
| 543 | + fitted = pm.fit(100, method="fullrank_advi", progressbar=False, random_seed=42) |
| 544 | + |
| 545 | + s = fitted.state |
| 546 | + assert s.mean.keys() == {"sigma", "mu"} |
| 547 | + assert s.std is not None |
| 548 | + assert s.mean["sigma"].values > 0 |
| 549 | + |
| 550 | + def test_state_empirical_std_is_none(self): |
| 551 | + """Empirical state has std=None.""" |
| 552 | + rng = np.random.default_rng(42) |
| 553 | + with pm.Model(): |
| 554 | + pm.Normal("mu", 0, 1) |
| 555 | + pm.Normal("y", rng.normal(size=10), observed=rng.normal(size=10)) |
| 556 | + inference = pm.SVGD(n_particles=50, random_seed=42) |
| 557 | + fitted = inference.fit(100, progressbar=False) |
| 558 | + |
| 559 | + s = fitted.state |
| 560 | + assert s.std is None |
| 561 | + assert "mu" in s.mean |
| 562 | + |
| 563 | + def test_state_is_single_group_approx_attr(self): |
| 564 | + """state is accessible from SingleGroupApproximation via __getattr__ proxy.""" |
| 565 | + with pm.Model(): |
| 566 | + pm.Normal("mu", 0, 1) |
| 567 | + inference = pm.ADVI(random_seed=42) |
| 568 | + fitted = inference.fit(10, progressbar=False) |
| 569 | + |
| 570 | + s = fitted.state |
| 571 | + assert "mu" in s.mean |
| 572 | + |
| 573 | + def test_state_in_callback(self): |
| 574 | + """Callbacks can access state during training.""" |
| 575 | + rng = np.random.default_rng(42) |
| 576 | + snapshots = [] |
| 577 | + |
| 578 | + def callback(approx, losses, i): |
| 579 | + s = approx.state |
| 580 | + snapshots.append( |
| 581 | + { |
| 582 | + "i": i, |
| 583 | + "mean": s.mean, |
| 584 | + "std": s.std, |
| 585 | + } |
| 586 | + ) |
| 587 | + |
| 588 | + with pm.Model(): |
| 589 | + pm.HalfNormal("sigma", sigma=5.0) |
| 590 | + pm.Normal("mu", 0, 1) |
| 591 | + pm.Normal("y", rng.normal(size=3), observed=rng.normal(size=3)) |
| 592 | + inference = pm.ADVI(random_seed=42) |
| 593 | + fitted = inference.fit(50, callbacks=[callback], progressbar=False) |
| 594 | + |
| 595 | + assert len(snapshots) == 50 |
| 596 | + for snap in snapshots: |
| 597 | + assert isinstance(snap["mean"], Dataset) |
| 598 | + assert set(snap["mean"].keys()) == {"sigma", "mu"} |
| 599 | + assert snap["std"] is not None |
| 600 | + assert set(snap["std"].keys()) == {"sigma", "mu"} |
| 601 | + # The last snapshot should match the final state |
| 602 | + final = fitted.state |
| 603 | + np.testing.assert_allclose( |
| 604 | + snapshots[-1]["mean"]["sigma"].values, final.mean["sigma"].values |
| 605 | + ) |
| 606 | + # Parameters should have moved from their initial values |
| 607 | + first_mean = snapshots[0]["mean"]["mu"].values |
| 608 | + last_mean = snapshots[-1]["mean"]["mu"].values |
| 609 | + assert not np.allclose(first_mean, last_mean), "parameters should change during training" |
| 610 | + |
| 611 | + def test_state_dirichlet(self): |
| 612 | + """State works with Dirichlet (simplex transform changes dimensionality).""" |
| 613 | + with pm.Model(): |
| 614 | + pm.Dirichlet("p", a=[1, 2, 3]) |
| 615 | + fitted = pm.fit(50, method="advi", progressbar=False, random_seed=42) |
| 616 | + |
| 617 | + s = fitted.state |
| 618 | + # Dirichlet with K=3 has K-1=2 unconstrained dims, 3 constrained dims |
| 619 | + assert "p" in s.mean |
| 620 | + assert s.mean["p"].values.shape == (3,) |
| 621 | + # Values should be on the simplex (sum to 1) |
| 622 | + np.testing.assert_allclose(s.mean["p"].values.sum(), 1.0, atol=1e-6) |
| 623 | + assert (s.mean["p"].values >= 0).all() |
| 624 | + assert (s.mean["p"].values <= 1).all() |
| 625 | + # std should also be in constrained space |
| 626 | + assert s.std is not None |
| 627 | + assert "p" in s.std |
| 628 | + assert s.std["p"].values.shape == (3,) |
| 629 | + |
| 630 | + def test_state_include_transformed(self): |
| 631 | + """include_transformed=True adds unconstrained variables to state.""" |
| 632 | + with pm.Model(): |
| 633 | + pm.HalfNormal("sigma", sigma=5.0) |
| 634 | + pm.Normal("mu", 0, 1) |
| 635 | + fitted = pm.fit( |
| 636 | + 50, |
| 637 | + method="advi", |
| 638 | + progressbar=False, |
| 639 | + random_seed=42, |
| 640 | + include_transformed=True, |
| 641 | + ) |
| 642 | + |
| 643 | + s = fitted.state |
| 644 | + # Constrained variables always present |
| 645 | + assert "sigma" in s.mean |
| 646 | + assert "mu" in s.mean |
| 647 | + # Unconstrained variables included when include_transformed=True |
| 648 | + assert "sigma_log__" in s.mean |
| 649 | + # mu has no transform, so it appears only once |
| 650 | + assert list(s.mean.data_vars) == ["sigma", "mu", "sigma_log__"] |
| 651 | + assert s.std is not None |
| 652 | + assert "sigma_log__" in s.std |
| 653 | + |
| 654 | + def test_state_include_transformed_dirichlet(self): |
| 655 | + """include_transformed=True with Dirichlet (dimensionality-changing transform).""" |
| 656 | + with pm.Model(): |
| 657 | + pm.Dirichlet("p", a=[1, 2, 3]) |
| 658 | + fitted = pm.fit( |
| 659 | + 50, |
| 660 | + method="advi", |
| 661 | + progressbar=False, |
| 662 | + random_seed=42, |
| 663 | + include_transformed=True, |
| 664 | + ) |
| 665 | + |
| 666 | + s = fitted.state |
| 667 | + # Constrained: p (shape 3, on simplex) |
| 668 | + assert "p" in s.mean |
| 669 | + assert s.mean["p"].values.shape == (3,) |
| 670 | + np.testing.assert_allclose(s.mean["p"].values.sum(), 1.0, atol=1e-6) |
| 671 | + # Unconstrained: p_simplex__ (shape 2, K-1 dims) |
| 672 | + assert "p_simplex__" in s.mean |
| 673 | + assert s.mean["p_simplex__"].values.shape == (2,) |
0 commit comments