diff --git a/CHANGELOG.md b/CHANGELOG.md index 6fa302b..cfa8a91 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -13,6 +13,10 @@ after its first published distribution. - Public API stability policy for tier-1 modules. - README Python API example coverage through a contract test. - Changelog discipline before first publish. +- Internal masked uncertainty-weighted multitask loss for auxiliary training + heads, combining only the targets each sample actually carries. +- Internal hierarchical utterance sampling primitives: square-root corpus mass, + inverse-square-root class mass, and bounded seeded window selection. ### Changed diff --git a/ser/_internal/heads/multitask_loss.py b/ser/_internal/heads/multitask_loss.py new file mode 100644 index 0000000..babb202 --- /dev/null +++ b/ser/_internal/heads/multitask_loss.py @@ -0,0 +1,65 @@ +"""Masked uncertainty-weighted objectives for auxiliary training heads.""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence + +import torch +from torch import Tensor, nn + + +class MaskedUncertaintyWeightedLoss(nn.Module): + """Combines available per-sample task losses with learned uncertainty weights.""" + + def __init__( + self, + tasks: Sequence[str], + *, + primary_task: str = "primary_emotion", + minimum_primary_weight: float = 0.25, + ) -> None: + """Initializes trainable log variances for a fixed task set.""" + super().__init__() + normalized_tasks = tuple(dict.fromkeys(task.strip() for task in tasks if task.strip())) + if not normalized_tasks: + raise ValueError("At least one multitask objective is required.") + if not 0.0 < minimum_primary_weight <= 1.0: + raise ValueError("minimum_primary_weight must be within (0, 1].") + if any("." in task for task in normalized_tasks): + raise ValueError("Task names cannot contain '.'.") + self.primary_task = primary_task + self.minimum_primary_weight = minimum_primary_weight + self.log_variances = nn.ParameterDict( + {task: nn.Parameter(torch.zeros((), dtype=torch.float32)) for task in normalized_tasks} + ) + + def forward( + self, + losses: Mapping[str, Tensor], + masks: Mapping[str, Tensor], + ) -> Tensor: + """Returns a scalar loss using only targets marked available by each mask.""" + total: Tensor | None = None + active_tasks = 0 + for task, log_variance in self.log_variances.items(): + if task not in losses or task not in masks: + continue + task_losses = losses[task] + mask = masks[task] + if task_losses.ndim == 0: + task_losses = task_losses.unsqueeze(0) + if mask.shape != task_losses.shape: + raise ValueError(f"Loss and mask shapes differ for task {task!r}.") + active = mask.to(dtype=torch.bool) + if not bool(torch.any(active)): + continue + mean_loss = task_losses[active].mean() + weight = torch.exp(-log_variance) + if task == self.primary_task: + weight = torch.clamp_min(weight, self.minimum_primary_weight) + weighted = weight * mean_loss + log_variance + total = weighted if total is None else total + weighted + active_tasks += 1 + if total is None or active_tasks == 0: + raise ValueError("No available targets were supplied to the multitask loss.") + return total diff --git a/ser/_internal/models/utterance_sampling.py b/ser/_internal/models/utterance_sampling.py new file mode 100644 index 0000000..612b48a --- /dev/null +++ b/ser/_internal/models/utterance_sampling.py @@ -0,0 +1,120 @@ +"""Utterance-level corpus/class sampling and bounded window selection.""" + +from __future__ import annotations + +import hashlib +import math +import random +from collections import Counter, defaultdict +from dataclasses import dataclass + + +@dataclass(frozen=True) +class UtteranceSamplingItem: + """Minimal utterance metadata needed by the balanced sampler.""" + + sample_id: str + corpus: str + label: str + window_count: int + duration_seconds: float | None = None + + def validate(self) -> None: + """Validates item identity and bounded integer window count.""" + if not self.sample_id.strip() or not self.corpus.strip() or not self.label.strip(): + raise ValueError("Sampling item identifiers and label must be non-empty.") + if self.window_count <= 0: + raise ValueError("Sampling item window_count must be positive.") + if self.duration_seconds is not None and self.duration_seconds <= 0.0: + raise ValueError("Sampling item duration_seconds must be positive when provided.") + + +@dataclass(frozen=True) +class SamplingProbability: + """Expected contribution of one utterance under hierarchical sampling.""" + + sample_id: str + corpus: str + label: str + probability: float + + +def utterance_sampling_distribution( + items: list[UtteranceSamplingItem], +) -> tuple[SamplingProbability, ...]: + """Computes ``sqrt(corpus)`` and inverse-``sqrt(class)`` sampling probabilities.""" + if not items: + raise ValueError("Cannot build a sampling distribution for an empty dataset.") + sample_ids: set[str] = set() + corpus_counts: Counter[str] = Counter() + class_counts: dict[str, Counter[str]] = defaultdict(Counter) + for item in items: + item.validate() + if item.sample_id in sample_ids: + raise ValueError(f"Duplicate sampling sample_id {item.sample_id!r}.") + sample_ids.add(item.sample_id) + corpus_counts[item.corpus] += 1 + class_counts[item.corpus][item.label] += 1 + + corpus_normalizer = sum(math.sqrt(count) for count in corpus_counts.values()) + class_normalizers = { + corpus: sum(1.0 / math.sqrt(count) for count in counts.values()) + for corpus, counts in class_counts.items() + } + probabilities = [] + for item in items: + corpus_probability = math.sqrt(corpus_counts[item.corpus]) / corpus_normalizer + label_count = class_counts[item.corpus][item.label] + class_probability = (1.0 / math.sqrt(label_count)) / class_normalizers[item.corpus] + item_probability = corpus_probability * class_probability / label_count + probabilities.append( + SamplingProbability(item.sample_id, item.corpus, item.label, item_probability) + ) + total = sum(row.probability for row in probabilities) + if not math.isclose(total, 1.0, rel_tol=1e-12, abs_tol=1e-12): + raise RuntimeError(f"Sampling probabilities do not sum to one: {total!r}.") + return tuple(sorted(probabilities, key=lambda row: row.sample_id)) + + +def select_training_windows( + *, + sample_id: str, + window_count: int, + max_windows: int, + seed: int, + epoch: int = 0, +) -> tuple[int, ...]: + """Selects a deterministic random bounded window subset for one epoch.""" + if not sample_id.strip(): + raise ValueError("sample_id must be non-empty.") + if window_count <= 0 or max_windows <= 0: + raise ValueError("window_count and max_windows must be positive.") + if epoch < 0: + raise ValueError("epoch must be non-negative.") + if window_count <= max_windows: + return tuple(range(window_count)) + digest = hashlib.sha256(f"{seed}:{epoch}:{sample_id}".encode()).digest() + rng = random.Random(int.from_bytes(digest[:8], "big")) + return tuple(sorted(rng.sample(range(window_count), max_windows))) + + +def sampling_contributions( + items: list[UtteranceSamplingItem], +) -> dict[str, dict[str, float]]: + """Reports expected sample and duration contributions by corpus and class.""" + item_by_id = {item.sample_id: item for item in items} + probabilities = utterance_sampling_distribution(items) + corpus: defaultdict[str, float] = defaultdict(float) + classes: defaultdict[str, float] = defaultdict(float) + duration: defaultdict[str, float] = defaultdict(float) + for row in probabilities: + corpus[row.corpus] += row.probability + classes[f"{row.corpus}:{row.label}"] += row.probability + seconds = item_by_id[row.sample_id].duration_seconds + if seconds is not None: + duration[row.corpus] += row.probability * seconds + return { + "corpus": dict(sorted(corpus.items())), + "class": dict(sorted(classes.items())), + "expected_duration_seconds": dict(sorted(duration.items())), + } diff --git a/tests/suites/unit/heads/test_multitask_loss.py b/tests/suites/unit/heads/test_multitask_loss.py new file mode 100644 index 0000000..29a2b08 --- /dev/null +++ b/tests/suites/unit/heads/test_multitask_loss.py @@ -0,0 +1,40 @@ +"""Tests for masked uncertainty-weighted multitask losses.""" + +from __future__ import annotations + +import pytest +import torch + +from ser._internal.heads.multitask_loss import MaskedUncertaintyWeightedLoss + + +def test_missing_auxiliary_targets_contribute_no_loss_or_gradient() -> None: + """Unavailable auxiliary labels remain isolated from the primary objective.""" + objective = MaskedUncertaintyWeightedLoss(("primary_emotion", "vad")) + primary = torch.tensor([1.0, 3.0], requires_grad=True) + auxiliary = torch.tensor([100.0, 200.0], requires_grad=True) + + value = objective( + {"primary_emotion": primary, "vad": auxiliary}, + { + "primary_emotion": torch.tensor([True, True]), + "vad": torch.tensor([False, False]), + }, + ) + value.backward() + + assert value.item() == pytest.approx(2.0) + assert primary.grad is not None + assert auxiliary.grad is None + assert objective.log_variances["vad"].grad is None + + +def test_loss_rejects_batches_without_any_available_target() -> None: + """Silent zero-loss batches fail closed.""" + objective = MaskedUncertaintyWeightedLoss(("primary_emotion",)) + + with pytest.raises(ValueError, match="No available targets"): + objective( + {"primary_emotion": torch.tensor([1.0])}, + {"primary_emotion": torch.tensor([False])}, + ) diff --git a/tests/suites/unit/models/test_utterance_sampling.py b/tests/suites/unit/models/test_utterance_sampling.py new file mode 100644 index 0000000..1a8819b --- /dev/null +++ b/tests/suites/unit/models/test_utterance_sampling.py @@ -0,0 +1,48 @@ +"""Tests for corpus/class balanced utterance sampling.""" + +from __future__ import annotations + +import math + +from ser._internal.models.utterance_sampling import ( + UtteranceSamplingItem, + select_training_windows, + utterance_sampling_distribution, +) + + +def test_distribution_uses_sqrt_corpus_and_inverse_sqrt_class_weights() -> None: + """Expected hierarchical mass matches the declared sampling policy.""" + items = [ + UtteranceSamplingItem("a:happy:1", "a", "happy", 10), + UtteranceSamplingItem("a:happy:2", "a", "happy", 20), + UtteranceSamplingItem("a:sad:1", "a", "sad", 1), + UtteranceSamplingItem("b:sad:1", "b", "sad", 1), + ] + rows = utterance_sampling_distribution(items) + corpus_a = sum(row.probability for row in rows if row.corpus == "a") + corpus_b = sum(row.probability for row in rows if row.corpus == "b") + expected_a = math.sqrt(3) / (math.sqrt(3) + 1) + + assert math.isclose(corpus_a, expected_a) + assert math.isclose(corpus_b, 1.0 - expected_a) + happy_mass = sum(row.probability for row in rows if row.corpus == "a" and row.label == "happy") + sad_mass = sum(row.probability for row in rows if row.corpus == "a" and row.label == "sad") + assert sad_mass > happy_mass + + +def test_training_window_selection_is_bounded_seeded_and_epoch_varying() -> None: + """Long utterances never materialize an unbounded training-window contribution.""" + first = select_training_windows( + sample_id="corpus:1", window_count=100, max_windows=4, seed=7, epoch=0 + ) + repeated = select_training_windows( + sample_id="corpus:1", window_count=100, max_windows=4, seed=7, epoch=0 + ) + next_epoch = select_training_windows( + sample_id="corpus:1", window_count=100, max_windows=4, seed=7, epoch=1 + ) + + assert len(first) == 4 + assert first == repeated + assert first != next_epoch