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
2 changes: 2 additions & 0 deletions docs/source/api/metrics.rst
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,8 @@ Generalizability

.. autofunction:: torch_measure.metrics.g_coefficient

.. autofunction:: torch_measure.metrics.intraclass_correlation

.. autofunction:: torch_measure.metrics.d_study

Calibration
Expand Down
8 changes: 7 additions & 1 deletion src/torch_measure/metrics/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,12 @@

from torch_measure.metrics.calibration import brier_score, expected_calibration_error
from torch_measure.metrics.correlation import point_biserial_correlation, tetrachoric_correlation
from torch_measure.metrics.generalizability import d_study, g_coefficient, variance_components
from torch_measure.metrics.generalizability import (
d_study,
g_coefficient,
intraclass_correlation,
variance_components,
)
from torch_measure.metrics.network import (
betweenness_centrality,
closeness_centrality,
Expand Down Expand Up @@ -34,6 +39,7 @@
"cronbach_alpha",
"variance_components",
"g_coefficient",
"intraclass_correlation",
"d_study",
"mokken_scalability",
"expected_calibration_error",
Expand Down
61 changes: 61 additions & 0 deletions src/torch_measure/metrics/generalizability.py
Original file line number Diff line number Diff line change
Expand Up @@ -175,6 +175,67 @@ def g_coefficient(
return s_p / denom


def intraclass_correlation(
variance_components: dict,
form: str = "ICC3k",
n_items: int | None = None,
) -> float:
"""Intraclass correlation coefficient from two-way variance components.

Subjects are targets, items are raters. ICC2/ICC3 are single-rater
(absolute agreement / consistency); ICC2k/ICC3k average over k raters and
equal the absolute / relative :func:`g_coefficient` at ``n_reps=1``.
One-way forms (ICC1) need a one-way model and are not supported here.

Parameters
----------
variance_components : dict
Output of :func:`variance_components`, or any dict with keys
``subject``, ``item``, ``subject_item``, ``residual``.
form : {"ICC2", "ICC3", "ICC2k", "ICC3k"}
Which coefficient to compute. The ``k`` forms average over raters.
n_items : int | None
Number of raters k for the ``k`` forms; defaults to
``variance_components["n_items"]``.

Returns
-------
float
ICC in [0, 1]. 0.0 if the denominator is numerically zero.
"""
required = {"subject", "item", "subject_item", "residual"}
missing = required - set(variance_components)
if missing:
raise ValueError(f"Missing required keys: {sorted(missing)}.")
if form in {"ICC1", "ICC1k"}:
raise ValueError(f"{form} requires a one-way model and is not supported; use ICC2/ICC3/ICC2k/ICC3k.")
if form not in {"ICC2", "ICC3", "ICC2k", "ICC3k"}:
raise ValueError(f"Unknown form: {form!r}. Expected one of ICC2, ICC3, ICC2k, ICC3k.")

s_p = float(variance_components["subject"])
s_i = float(variance_components["item"])
s_pi = float(variance_components["subject_item"])
s_e = float(variance_components["residual"])

averaged = form.endswith("k")
absolute = form in {"ICC2", "ICC2k"}

if averaged:
k = n_items if n_items is not None else int(variance_components["n_items"])
if k < 1:
raise ValueError(f"n_items must be >= 1; got {k}.")
else:
k = 1

err = (s_i + s_pi + s_e) if absolute else (s_pi + s_e)
err = err / k

denom = s_p + err
if denom < 1e-12:
return 0.0
return s_p / denom


def d_study(
variance_components: dict,
n_items_grid: Sequence[int],
Expand Down
65 changes: 65 additions & 0 deletions tests/test_metrics/test_generalizability.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
from torch_measure.metrics.generalizability import (
d_study,
g_coefficient,
intraclass_correlation,
variance_components,
)

Expand Down Expand Up @@ -214,6 +215,70 @@ def test_empty_grid_raises(self):
d_study(self._vc(), n_items_grid=[], n_reps_grid=[1])


class TestIntraclassCorrelation:
def _vc(self) -> dict:
return {
"subject": 1.0,
"item": 0.5,
"subject_item": 0.3,
"residual": 0.2,
"n_items": 10,
}

def test_in_unit_interval(self):
for form in ("ICC2", "ICC3", "ICC2k", "ICC3k"):
icc = intraclass_correlation(self._vc(), form=form)
assert 0.0 <= icc <= 1.0

def test_equivalence_with_g_coefficient(self):
vc = self._vc()
k = 8
assert intraclass_correlation(vc, "ICC2k", n_items=k) == pytest.approx(
g_coefficient(vc, n_items=k, n_reps=1, type="absolute")
)
assert intraclass_correlation(vc, "ICC3k", n_items=k) == pytest.approx(
g_coefficient(vc, n_items=k, n_reps=1, type="relative")
)

def test_consistency_ge_absolute(self):
vc = self._vc()
assert intraclass_correlation(vc, "ICC3") >= intraclass_correlation(vc, "ICC2")
assert intraclass_correlation(vc, "ICC3k") >= intraclass_correlation(vc, "ICC2k")

def test_average_ge_single(self):
vc = self._vc()
assert intraclass_correlation(vc, "ICC3k", n_items=10) >= intraclass_correlation(vc, "ICC3")
assert intraclass_correlation(vc, "ICC2k", n_items=10) >= intraclass_correlation(vc, "ICC2")

def test_defaults_n_items_from_dict(self):
vc = self._vc()
assert intraclass_correlation(vc, "ICC3k") == pytest.approx(
intraclass_correlation(vc, "ICC3k", n_items=vc["n_items"])
)

def test_icc1_raises(self):
with pytest.raises(ValueError, match="one-way model"):
intraclass_correlation(self._vc(), form="ICC1")

def test_unknown_form_raises(self):
with pytest.raises(ValueError, match="Unknown form"):
intraclass_correlation(self._vc(), form="bogus")

def test_missing_keys_raises(self):
with pytest.raises(ValueError, match="Missing required keys"):
intraclass_correlation({"subject": 1.0}, form="ICC3")

def test_zero_components_returns_zero(self):
vc = {"subject": 0.0, "item": 0.0, "subject_item": 0.0, "residual": 0.0, "n_items": 5}
assert intraclass_correlation(vc, "ICC2") == 0.0

def test_from_real_variance_components(self):
df = _synth_crossed_design(n_p=60, n_i=12, n_r=2, seed=1)
vc = variance_components(df)
icc3k = intraclass_correlation(vc, "ICC3k")
assert icc3k == pytest.approx(g_coefficient(vc, n_items=vc["n_items"], n_reps=1, type="relative"))


def test_end_to_end_pipeline():
"""variance_components -> g_coefficient -> d_study composes cleanly."""
df = _synth_crossed_design(n_p=40, n_i=15, n_r=3, seed=0)
Expand Down
Loading