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: 4 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
65 changes: 65 additions & 0 deletions ser/_internal/heads/multitask_loss.py
Original file line number Diff line number Diff line change
@@ -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
120 changes: 120 additions & 0 deletions ser/_internal/models/utterance_sampling.py
Original file line number Diff line number Diff line change
@@ -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())),
}
40 changes: 40 additions & 0 deletions tests/suites/unit/heads/test_multitask_loss.py
Original file line number Diff line number Diff line change
@@ -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])},
)
48 changes: 48 additions & 0 deletions tests/suites/unit/models/test_utterance_sampling.py
Original file line number Diff line number Diff line change
@@ -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