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 birdman/default_models.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,13 @@
from os.path import join as pjoin
from pkg_resources import resource_filename
from pathlib import Path

import biom
import numpy as np
import pandas as pd

from .model_base import TableModel, SingleFeatureModel

TEMPLATES = resource_filename("birdman", "templates")
TEMPLATES = str(Path(__file__).parent / "templates")
DEFAULT_MODEL_DICT = {
"negative_binomial": {
"standard": pjoin(TEMPLATES, "negative_binomial.stan"),
Expand Down
6 changes: 3 additions & 3 deletions birdman/transform.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,11 +28,11 @@ def posterior_alr_to_clr(
new_posterior = posterior.copy()
for param in alr_params:
param_da = posterior[param]
all_chain_alr_coords = param_da
all_chain_clr_coords = []

for i, chain_alr_coords in all_chain_alr_coords.groupby("chain"):
chain_clr_coords = _beta_alr_to_clr(chain_alr_coords)
for chain_idx in range(param_da.sizes["chain"]):
chain_data = param_da.isel(chain=chain_idx).values
chain_clr_coords = _beta_alr_to_clr(chain_data)
all_chain_clr_coords.append(chain_clr_coords)

all_chain_clr_coords = np.array(all_chain_clr_coords)
Expand Down
4 changes: 2 additions & 2 deletions tests/conftest.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,13 @@
import os
from pkg_resources import resource_filename
from pathlib import Path

import biom
import pandas as pd
import pytest

from birdman import NegativeBinomial, NegativeBinomialSingle

TEST_DATA = resource_filename("tests", "data")
TEST_DATA = str(Path(__file__).parent / "data")
TBL_FILE = os.path.join(TEST_DATA, "macaque_tbl.biom")
MD_FILE = os.path.join(TEST_DATA, "macaque_metadata.tsv")

Expand Down
4 changes: 2 additions & 2 deletions tests/test_custom_model.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from pkg_resources import resource_filename
from pathlib import Path

import numpy as np

Expand All @@ -13,7 +13,7 @@ def test_custom_model(table_biom, metadata):

custom_model = TableModel(
table=table_biom,
model_path=resource_filename("tests", "custom_model.stan"),
model_path=str(Path(__file__).parent / "custom_model.stan"),
)
custom_model.create_regression(
formula="host_common_name",
Expand Down
6 changes: 4 additions & 2 deletions tests/test_model.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,15 @@
import os
from pkg_resources import resource_filename
from pathlib import Path

import numpy as np

from birdman import (NegativeBinomial, NegativeBinomialLME,
NegativeBinomialSingle, NegativeBinomialLMESingle,
ModelIterator)

TEMPLATES = resource_filename("birdman", "templates")
TEMPLATES = str(
Path(__file__).resolve().parent.parent / "birdman" / "templates"
)


class TestModelInheritance:
Expand Down
132 changes: 132 additions & 0 deletions tests/test_transform.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import numpy as np
import xarray as xr
from skbio.stats.composition import alr, clr

from birdman import transform
Expand Down Expand Up @@ -80,3 +81,134 @@ def test_convert_beta_coordinates():
clr_coords_sums = clr_coords.sum(axis=2)
exp_clr_coords_sums = np.zeros((2, 4))
np.testing.assert_array_almost_equal(exp_clr_coords_sums, clr_coords_sums)


def test_posterior_alr_to_clr_multi_chain():
"""Regression test: groupby chain dimension not squeezed."""
num_chains, num_draws, num_covariates, num_features_alr = 4, 50, 3, 5
num_features = num_features_alr + 1

beta_data = np.random.randn(
num_chains, num_draws, num_covariates, num_features_alr
)
feature_names = [f"feat{i}" for i in range(num_features)]
feature_names_alr = feature_names[1:]
covariate_names = [f"cov{i}" for i in range(num_covariates)]

ds = xr.Dataset({
"beta_var": xr.DataArray(
beta_data,
dims=["chain", "draw", "covariate", "feature_alr"],
coords={
"chain": np.arange(num_chains),
"draw": np.arange(num_draws),
"covariate": covariate_names,
"feature_alr": feature_names_alr,
},
)
})

result = transform.posterior_alr_to_clr(
ds,
alr_params=["beta_var"],
dim_replacement={"feature_alr": "feature"},
new_labels=feature_names,
)

assert set(result.dims) == {"chain", "draw", "feature", "covariate"}
assert result["beta_var"].shape == (
num_chains, num_draws, num_covariates, num_features
)
np.testing.assert_equal(result.coords["feature"].values, feature_names)
assert np.allclose(result["beta_var"].values.sum(axis=-1), 0, atol=1e-10)


def test_posterior_alr_to_clr_single_chain_single_covariate():
"""Verify fix works with VI (1 chain) and intercept-only (1 covariate)."""
num_chains, num_draws, num_covariates, num_features_alr = 1, 50, 1, 5
num_features = num_features_alr + 1

beta_data = np.random.randn(
num_chains, num_draws, num_covariates, num_features_alr
)
feature_names = [f"feat{i}" for i in range(num_features)]
feature_names_alr = feature_names[1:]
covariate_names = ["Intercept"]

ds = xr.Dataset({
"beta_var": xr.DataArray(
beta_data,
dims=["chain", "draw", "covariate", "feature_alr"],
coords={
"chain": np.arange(num_chains),
"draw": np.arange(num_draws),
"covariate": covariate_names,
"feature_alr": feature_names_alr,
},
)
})

result = transform.posterior_alr_to_clr(
ds,
alr_params=["beta_var"],
dim_replacement={"feature_alr": "feature"},
new_labels=feature_names,
)

assert result["beta_var"].shape == (
num_chains, num_draws, num_covariates, num_features
)
assert np.allclose(result["beta_var"].values.sum(axis=-1), 0, atol=1e-10)


def test_posterior_alr_to_clr_multiple_params():
"""Verify multiple alr_params (e.g. beta_var + subj_int) all transform."""
num_chains, num_draws = 2, 50
num_covariates, num_groups = 3, 4
num_features_alr, num_features = 5, 6

feature_names = [f"feat{i}" for i in range(num_features)]
feature_names_alr = feature_names[1:]

ds = xr.Dataset({
"beta_var": xr.DataArray(
np.random.randn(
num_chains, num_draws, num_covariates, num_features_alr
),
dims=["chain", "draw", "covariate", "feature_alr"],
coords={
"chain": np.arange(num_chains),
"draw": np.arange(num_draws),
"covariate": [f"cov{i}" for i in range(num_covariates)],
"feature_alr": feature_names_alr,
},
),
"subj_int": xr.DataArray(
np.random.randn(
num_chains, num_draws, num_groups, num_features_alr
),
dims=["chain", "draw", "group", "feature_alr"],
coords={
"chain": np.arange(num_chains),
"draw": np.arange(num_draws),
"group": [f"subj{i}" for i in range(num_groups)],
"feature_alr": feature_names_alr,
},
),
})

result = transform.posterior_alr_to_clr(
ds,
alr_params=["beta_var", "subj_int"],
dim_replacement={"feature_alr": "feature"},
new_labels=feature_names,
)

assert result["beta_var"].shape == (
num_chains, num_draws, num_covariates, num_features
)
assert result["subj_int"].shape == (
num_chains, num_draws, num_groups, num_features
)
assert np.allclose(result["beta_var"].values.sum(axis=-1), 0, atol=1e-10)
assert np.allclose(result["subj_int"].values.sum(axis=-1), 0, atol=1e-10)
Loading