Skip to content

geno_lewm.data.membership

membership

Explicit, checksum-bound variant membership for v0.3 datasets.

MEMBERSHIP_SCHEMA_VERSION module-attribute

MEMBERSHIP_SCHEMA_VERSION = '1.0.0'

Schema version for canonical membership artifacts and policies.

REQUIRED_MEMBERSHIP_ROLES module-attribute

REQUIRED_MEMBERSHIP_ROLES: tuple[str, ...] = ('train', 'validation', 'evaluation')

Roles that every membership contract must represent non-vacuously.

V03_CHROMOSOME_ROLES module-attribute

V03_CHROMOSOME_ROLES = ChromosomeRoles(train=(*(map(str, range(1, 20))), '22'), validation=('20',), evaluation=('21',))

Canonical chromosome split for the corrected v0.3 dataset.

ChromosomeRoles dataclass

ChromosomeRoles(train: tuple[str, ...], validation: tuple[str, ...], evaluation: tuple[str, ...])

Disjoint chromosome assignments for train, validation, and evaluation.

role_for

role_for(chrom: str) -> str

Return the explicit role for chrom or fail if it is unassigned.

Source code in geno_lewm/data/membership.py
def role_for(self, chrom: str) -> str:
    """Return the explicit role for ``chrom`` or fail if it is unassigned."""
    canonical = canonicalize_chromosome(chrom)
    for role in REQUIRED_MEMBERSHIP_ROLES:
        if canonical in getattr(self, role):
            return role
    raise InputError(
        "chromosome is not assigned to a membership role",
        details={"chrom": canonical},
    )

to_dict

to_dict() -> dict[str, list[str]]

Return the canonical JSON-native chromosome-role payload.

Source code in geno_lewm/data/membership.py
def to_dict(self) -> dict[str, list[str]]:
    """Return the canonical JSON-native chromosome-role payload."""
    return {role: list(getattr(self, role)) for role in REQUIRED_MEMBERSHIP_ROLES}

from_dict classmethod

from_dict(payload: Mapping[str, object]) -> ChromosomeRoles

Parse an explicit chromosome-role payload.

Source code in geno_lewm/data/membership.py
@classmethod
def from_dict(cls, payload: Mapping[str, object]) -> ChromosomeRoles:
    """Parse an explicit chromosome-role payload."""
    required = set(REQUIRED_MEMBERSHIP_ROLES)
    if not isinstance(payload, Mapping):
        raise InputError(
            "chromosome roles must be a mapping",
            details={"type": type(payload).__name__},
        )
    missing = required - set(payload)
    extra = set(payload) - required
    if missing or extra:
        raise InputError(
            "chromosome role keys do not match the schema",
            details={"missing": sorted(missing), "extra": sorted(extra)},
        )
    values: dict[str, tuple[str, ...]] = {}
    for role in REQUIRED_MEMBERSHIP_ROLES:
        raw = payload[role]
        if isinstance(raw, str | bytes) or not isinstance(raw, Sequence):
            raise InputError(
                f"{role} chromosomes must be a sequence",
                details={"role": role, "type": type(raw).__name__},
            )
        if not all(isinstance(value, str) for value in raw):
            raise InputError(
                f"{role} chromosomes must contain strings",
                details={"role": role},
            )
        values[role] = tuple(cast(Sequence[str], raw))
    return cls(
        train=values["train"],
        validation=values["validation"],
        evaluation=values["evaluation"],
    )

MembershipRow dataclass

MembershipRow(variant: CanonicalVariant, role: str, reason_mask: int, source: str, source_row_id: str)

One canonical variant assignment with source-level provenance.

to_dict

to_dict() -> dict[str, object]

Return the strict artifact-row payload.

Source code in geno_lewm/data/membership.py
def to_dict(self) -> dict[str, object]:
    """Return the strict artifact-row payload."""
    return {
        "variant_key": self.variant.key,
        "variant_digest": self.variant.digest,
        "role": self.role,
        "reason_mask": self.reason_mask,
        "source": self.source,
        "source_row_id": self.source_row_id,
    }

from_dict classmethod

from_dict(payload: Mapping[str, object]) -> MembershipRow

Parse one strict row and verify its key/digest binding.

Source code in geno_lewm/data/membership.py
@classmethod
def from_dict(cls, payload: Mapping[str, object]) -> MembershipRow:
    """Parse one strict row and verify its key/digest binding."""
    required = {
        "variant_key",
        "variant_digest",
        "role",
        "reason_mask",
        "source",
        "source_row_id",
    }
    if not isinstance(payload, Mapping):
        raise InputError(
            "membership row must be a mapping",
            details={"type": type(payload).__name__},
        )
    missing = required - set(payload)
    extra = set(payload) - required
    if missing or extra:
        raise InputError(
            "membership row keys do not match the schema",
            details={"missing": sorted(missing), "extra": sorted(extra)},
        )
    variant_key = payload["variant_key"]
    if not isinstance(variant_key, str):
        raise InputError(
            "membership variant_key must be a canonical variant key",
            details={"value": variant_key},
        )
    variant = CanonicalVariant.from_key(variant_key)
    if payload["variant_digest"] != variant.digest:
        raise InputError(
            "membership variant digest does not match its canonical key",
            details={
                "variant_key": variant.key,
                "declared": payload["variant_digest"],
                "computed": variant.digest,
            },
        )
    return cls(
        variant=variant,
        role=payload["role"],  # type: ignore[arg-type]
        reason_mask=payload["reason_mask"],  # type: ignore[arg-type]
        source=payload["source"],  # type: ignore[arg-type]
        source_row_id=payload["source_row_id"],  # type: ignore[arg-type]
    )

MembershipArtifact dataclass

MembershipArtifact(artifact_id: str, assembly: str, chromosome_roles: ChromosomeRoles, rows: tuple[MembershipRow, ...], schema_version: str = MEMBERSHIP_SCHEMA_VERSION)

Canonical membership rows plus the chromosome policy that assigned them.

role_counts property

role_counts: dict[str, int]

Return row counts for each required role.

row_count property

row_count: int

Return the number of source membership rows.

variant_count property

variant_count: int

Return the number of distinct canonical variants.

content_sha256 property

content_sha256: str

Return the SHA-256 digest of the canonical artifact payload.

to_dict

to_dict() -> dict[str, object]

Return the deterministic, checksum-addressed artifact payload.

Source code in geno_lewm/data/membership.py
def to_dict(self) -> dict[str, object]:
    """Return the deterministic, checksum-addressed artifact payload."""
    return {
        "schema_version": self.schema_version,
        "artifact_id": self.artifact_id,
        "assembly": self.assembly,
        "chromosome_roles": self.chromosome_roles.to_dict(),
        "rows": [row.to_dict() for row in self.rows],
    }

from_dict classmethod

from_dict(payload: Mapping[str, object]) -> MembershipArtifact

Parse and validate a strict membership artifact payload.

Source code in geno_lewm/data/membership.py
@classmethod
def from_dict(cls, payload: Mapping[str, object]) -> MembershipArtifact:
    """Parse and validate a strict membership artifact payload."""
    required = {"schema_version", "artifact_id", "assembly", "chromosome_roles", "rows"}
    if not isinstance(payload, Mapping):
        raise InputError(
            "membership artifact must be a mapping",
            details={"type": type(payload).__name__},
        )
    missing = required - set(payload)
    extra = set(payload) - required
    if missing or extra:
        raise InputError(
            "membership artifact keys do not match the schema",
            details={"missing": sorted(missing), "extra": sorted(extra)},
        )
    roles_payload = payload["chromosome_roles"]
    rows_payload = payload["rows"]
    if not isinstance(roles_payload, Mapping):
        raise InputError("membership chromosome_roles must be a mapping")
    if isinstance(rows_payload, str | bytes) or not isinstance(rows_payload, Sequence):
        raise InputError("membership rows must be a sequence")
    rows: list[MembershipRow] = []
    for row_payload in rows_payload:
        if not isinstance(row_payload, Mapping):
            raise InputError("membership rows must contain mappings")
        rows.append(MembershipRow.from_dict(row_payload))
    return cls(
        artifact_id=payload["artifact_id"],  # type: ignore[arg-type]
        assembly=payload["assembly"],  # type: ignore[arg-type]
        chromosome_roles=ChromosomeRoles.from_dict(roles_payload),
        rows=tuple(rows),
        schema_version=payload["schema_version"],  # type: ignore[arg-type]
    )

MembershipArtifactBinding dataclass

MembershipArtifactBinding(artifact_id: str, sha256: str, row_count: int, variant_count: int, train_rows: int, validation_rows: int, evaluation_rows: int)

Hash and count commitments for one membership artifact.

from_artifact classmethod

from_artifact(artifact: MembershipArtifact) -> MembershipArtifactBinding

Build an immutable commitment from validated artifact content.

Source code in geno_lewm/data/membership.py
@classmethod
def from_artifact(cls, artifact: MembershipArtifact) -> MembershipArtifactBinding:
    """Build an immutable commitment from validated artifact content."""
    counts = artifact.role_counts
    return cls(
        artifact_id=artifact.artifact_id,
        sha256=artifact.content_sha256,
        row_count=artifact.row_count,
        variant_count=artifact.variant_count,
        train_rows=counts["train"],
        validation_rows=counts["validation"],
        evaluation_rows=counts["evaluation"],
    )

to_dict

to_dict() -> dict[str, object]

Return the canonical JSON-native binding payload.

Source code in geno_lewm/data/membership.py
def to_dict(self) -> dict[str, object]:
    """Return the canonical JSON-native binding payload."""
    return {
        "artifact_id": self.artifact_id,
        "sha256": self.sha256,
        "row_count": self.row_count,
        "variant_count": self.variant_count,
        "role_counts": {
            "train": self.train_rows,
            "validation": self.validation_rows,
            "evaluation": self.evaluation_rows,
        },
    }

from_dict classmethod

from_dict(payload: Mapping[str, object], *, artifact: MembershipArtifact) -> MembershipArtifactBinding

Parse a strict binding and verify it against its source artifact.

Source code in geno_lewm/data/membership.py
@classmethod
def from_dict(
    cls,
    payload: Mapping[str, object],
    *,
    artifact: MembershipArtifact,
) -> MembershipArtifactBinding:
    """Parse a strict binding and verify it against its source artifact."""
    required = {"artifact_id", "sha256", "row_count", "variant_count", "role_counts"}
    if not isinstance(payload, Mapping):
        raise InputError(
            "membership artifact binding must be a mapping",
            details={"type": type(payload).__name__},
        )
    missing = required - set(payload)
    extra = set(payload) - required
    if missing or extra:
        raise InputError(
            "membership artifact binding keys do not match the schema",
            details={"missing": sorted(missing), "extra": sorted(extra)},
        )
    role_counts = payload["role_counts"]
    if not isinstance(role_counts, Mapping):
        raise InputError("membership artifact binding role_counts must be a mapping")
    role_missing = set(REQUIRED_MEMBERSHIP_ROLES) - set(role_counts)
    role_extra = set(role_counts) - set(REQUIRED_MEMBERSHIP_ROLES)
    if role_missing or role_extra:
        raise InputError(
            "membership artifact binding role_counts keys do not match the schema",
            details={"missing": sorted(role_missing), "extra": sorted(role_extra)},
        )
    binding = cls(
        artifact_id=payload["artifact_id"],  # type: ignore[arg-type]
        sha256=payload["sha256"],  # type: ignore[arg-type]
        row_count=payload["row_count"],  # type: ignore[arg-type]
        variant_count=payload["variant_count"],  # type: ignore[arg-type]
        train_rows=role_counts["train"],
        validation_rows=role_counts["validation"],
        evaluation_rows=role_counts["evaluation"],
    )
    if not isinstance(artifact, MembershipArtifact):
        raise InputError("binding artifact must be a MembershipArtifact")
    expected = cls.from_artifact(artifact)
    if binding != expected:
        raise InputError(
            "membership artifact binding does not match its source artifact",
            details={
                "artifact_id": artifact.artifact_id,
                "declared": binding.to_dict(),
                "expected": expected.to_dict(),
            },
        )
    return binding

MembershipHoldoutPolicy dataclass

MembershipHoldoutPolicy(assembly: str, chromosome_roles: ChromosomeRoles, artifact_bindings: tuple[MembershipArtifactBinding, ...], excluded_chromosomes: tuple[str, ...], excluded_variant_keys: tuple[str, ...], schema_version: str = MEMBERSHIP_SCHEMA_VERSION)

Validation/evaluation exclusions bound to membership hashes and counts.

identity property

identity: str

Return the canonical SHA-256 identity of this holdout policy.

excludes_variant

excludes_variant(variant: CanonicalVariant) -> bool

Return whether a canonical variant is withheld from training.

Source code in geno_lewm/data/membership.py
def excludes_variant(self, variant: CanonicalVariant) -> bool:
    """Return whether a canonical variant is withheld from training."""
    if not isinstance(variant, CanonicalVariant):
        raise InputError("variant must be a CanonicalVariant")
    if variant.assembly != self.assembly:
        raise InputError(
            "variant assembly does not match the holdout policy",
            details={"variant": variant.assembly, "policy": self.assembly},
        )
    return (
        variant.chrom in self.excluded_chromosomes or variant.key in self.excluded_variant_keys
    )

to_dict

to_dict() -> dict[str, object]

Return the checksum- and count-bound policy plus its identity.

Source code in geno_lewm/data/membership.py
def to_dict(self) -> dict[str, object]:
    """Return the checksum- and count-bound policy plus its identity."""
    payload = self._identity_payload()
    payload["policy_identity"] = self.identity
    return payload

from_dict classmethod

from_dict(payload: Mapping[str, object], *, artifacts: Sequence[MembershipArtifact]) -> MembershipHoldoutPolicy

Parse a strict policy, verify its identity, and re-derive its bindings.

Source code in geno_lewm/data/membership.py
@classmethod
def from_dict(
    cls,
    payload: Mapping[str, object],
    *,
    artifacts: Sequence[MembershipArtifact],
) -> MembershipHoldoutPolicy:
    """Parse a strict policy, verify its identity, and re-derive its bindings."""
    required = {
        "schema_version",
        "assembly",
        "chromosome_roles",
        "artifact_bindings",
        "excluded_chromosomes",
        "excluded_variant_keys",
        "policy_identity",
    }
    if not isinstance(payload, Mapping):
        raise InputError(
            "membership holdout policy must be a mapping",
            details={"type": type(payload).__name__},
        )
    missing = required - set(payload)
    extra = set(payload) - required
    if missing or extra:
        raise InputError(
            "membership holdout policy keys do not match the schema",
            details={"missing": sorted(missing), "extra": sorted(extra)},
        )
    roles_payload = payload["chromosome_roles"]
    bindings_payload = payload["artifact_bindings"]
    excluded_chromosomes = payload["excluded_chromosomes"]
    excluded_variant_keys = payload["excluded_variant_keys"]
    declared_identity = payload["policy_identity"]
    if not isinstance(roles_payload, Mapping):
        raise InputError("membership holdout chromosome_roles must be a mapping")
    if isinstance(bindings_payload, str | bytes) or not isinstance(bindings_payload, Sequence):
        raise InputError("membership holdout artifact_bindings must be a sequence")
    artifacts_value: object = artifacts
    if isinstance(artifacts_value, MembershipArtifact) or not isinstance(
        artifacts_value, Sequence
    ):
        raise InputError("membership holdout source artifacts must be a sequence")
    source_artifacts = tuple(artifacts)
    if not source_artifacts or not all(
        isinstance(artifact, MembershipArtifact) for artifact in source_artifacts
    ):
        raise InputError(
            "membership holdout source artifacts must contain MembershipArtifact values"
        )
    artifacts_by_id = {artifact.artifact_id: artifact for artifact in source_artifacts}
    if len(artifacts_by_id) != len(source_artifacts):
        raise InputError("membership holdout source artifacts must have unique identifiers")
    bindings: list[MembershipArtifactBinding] = []
    for binding_payload in bindings_payload:
        if not isinstance(binding_payload, Mapping):
            raise InputError("membership holdout artifact_bindings must contain mappings")
        artifact_id = binding_payload.get("artifact_id")
        if not isinstance(artifact_id, str) or artifact_id not in artifacts_by_id:
            raise InputError(
                "membership holdout binding artifact_id is absent from source artifacts",
                details={"artifact_id": artifact_id},
            )
        bindings.append(
            MembershipArtifactBinding.from_dict(
                binding_payload,
                artifact=artifacts_by_id[artifact_id],
            )
        )
    for name, values in (
        ("excluded_chromosomes", excluded_chromosomes),
        ("excluded_variant_keys", excluded_variant_keys),
    ):
        if isinstance(values, str | bytes) or not isinstance(values, Sequence):
            raise InputError(f"membership holdout {name} must be a sequence")
        if not all(isinstance(value, str) for value in values):
            raise InputError(f"membership holdout {name} must contain strings")
    if not isinstance(declared_identity, str) or not looks_like_sha256(declared_identity):
        raise InputError(
            "membership holdout policy_identity must be 'sha256:<64hex>'",
            details={"policy_identity": declared_identity},
        )
    policy = cls(
        assembly=payload["assembly"],  # type: ignore[arg-type]
        chromosome_roles=ChromosomeRoles.from_dict(roles_payload),
        artifact_bindings=tuple(bindings),
        excluded_chromosomes=tuple(cast(Sequence[str], excluded_chromosomes)),
        excluded_variant_keys=tuple(cast(Sequence[str], excluded_variant_keys)),
        schema_version=payload["schema_version"],  # type: ignore[arg-type]
    )
    if declared_identity != policy.identity:
        raise InputError(
            "membership holdout policy identity drift",
            details={"declared": declared_identity, "computed": policy.identity},
        )
    declared_hashes = {
        binding.artifact_id: binding.sha256 for binding in policy.artifact_bindings
    }
    expected = derive_holdout_policy(source_artifacts, expected_sha256=declared_hashes)
    if policy != expected:
        raise InputError(
            "membership holdout policy does not match its source artifacts",
            details={"declared": policy.to_dict(), "expected": expected.to_dict()},
        )
    return policy

derive_holdout_policy

derive_holdout_policy(artifacts: Sequence[MembershipArtifact], *, expected_sha256: Mapping[str, str] | None = None) -> MembershipHoldoutPolicy

Derive validation/evaluation exclusions from validated membership artifacts.

When expected_sha256 is provided, its keys must exactly match the artifact identifiers and every declared checksum must match the artifact's canonical payload. Either way, the returned policy records the computed checksums and role/row counts.

Source code in geno_lewm/data/membership.py
def derive_holdout_policy(
    artifacts: Sequence[MembershipArtifact],
    *,
    expected_sha256: Mapping[str, str] | None = None,
) -> MembershipHoldoutPolicy:
    """Derive validation/evaluation exclusions from validated membership artifacts.

    When ``expected_sha256`` is provided, its keys must exactly match the
    artifact identifiers and every declared checksum must match the artifact's
    canonical payload.  Either way, the returned policy records the computed
    checksums and role/row counts.
    """
    if isinstance(artifacts, MembershipArtifact) or not isinstance(artifacts, Sequence):
        raise InputError("membership artifacts must be a sequence")
    normalized = tuple(artifacts)
    if not normalized or not all(isinstance(item, MembershipArtifact) for item in normalized):
        raise InputError("membership artifacts must contain at least one MembershipArtifact")
    artifact_ids = [artifact.artifact_id for artifact in normalized]
    if len(set(artifact_ids)) != len(artifact_ids):
        raise InputError("membership artifacts must have unique artifact_id values")

    if expected_sha256 is not None:
        if not isinstance(expected_sha256, Mapping):
            raise InputError(
                "expected membership checksums must be a mapping",
                details={"type": type(expected_sha256).__name__},
            )
        missing = set(artifact_ids) - set(expected_sha256)
        extra = set(expected_sha256) - set(artifact_ids)
        if missing or extra:
            raise InputError(
                "membership checksum keys do not match artifact identifiers",
                details={"missing": sorted(missing), "extra": sorted(extra)},
            )
        for artifact in normalized:
            declared = expected_sha256[artifact.artifact_id]
            if not looks_like_sha256(declared) or declared != artifact.content_sha256:
                raise InputError(
                    "membership artifact checksum mismatch",
                    details={
                        "artifact_id": artifact.artifact_id,
                        "declared": declared,
                        "computed": artifact.content_sha256,
                    },
                )

    reference = normalized[0]
    for artifact in normalized[1:]:
        if artifact.assembly != reference.assembly:
            raise InputError(
                "membership artifacts must use one assembly",
                details={
                    "reference": reference.assembly,
                    "artifact": artifact.assembly,
                    "artifact_id": artifact.artifact_id,
                },
            )
        if artifact.chromosome_roles != reference.chromosome_roles:
            raise InputError(
                "membership artifacts must use identical chromosome roles",
                details={"artifact_id": artifact.artifact_id},
            )

    excluded_keys = {
        row.variant.key
        for artifact in normalized
        for row in artifact.rows
        if row.role in {"validation", "evaluation"}
    }
    return MembershipHoldoutPolicy(
        assembly=reference.assembly,
        chromosome_roles=reference.chromosome_roles,
        artifact_bindings=tuple(
            MembershipArtifactBinding.from_artifact(artifact) for artifact in normalized
        ),
        excluded_chromosomes=reference.chromosome_roles.validation
        + reference.chromosome_roles.evaluation,
        excluded_variant_keys=tuple(sorted(excluded_keys)),
    )