#!/usr/bin/env python """Experiment 9-v2: research-upgraded KSL+CASL+NSL unified encoder. This is a stronger follow-up to the E9 focused baseline. It keeps the same KSL+CASL+NSL task protocol but adds metric-learning machinery inspired by recent SLR and low-resource recognition work: - ArcFace/CosFace task heads for angular class separation - supervised contrastive loss from the Exp8/Exp8-v2 loop - optional class-center loss to tighten within-class clusters - a motion-aware multi-scale temporal pose stem and conditional adapters The wrapper preserves Exp8-compatible result JSON files, so the normal KCN aggregator can compare it directly with E9.1-E9.4. """ from __future__ import annotations import argparse import math import random import sys from pathlib import Path from typing import Any, Optional, Sequence import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.data import DataLoader from tqdm.auto import tqdm ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) from experiments import exp8_unified_mixed_encoder as exp8 # noqa: E402 from experiments import exp8_v2_strong_unified_encoder as exp8_v2 # noqa: E402 BaseStrongUnifiedEncoder = exp8_v2.StrongUnifiedAfriSignEncoder class MotionAwarePoseStem(nn.Module): """Temporal pose stem with velocity, acceleration, and multi-scale motion filters.""" def __init__(self, feature_dim: int, hidden_dim: int, dropout: float) -> None: super().__init__() self.pose_proj = nn.Linear(feature_dim, hidden_dim) self.vel_proj = nn.Linear(feature_dim, hidden_dim) self.acc_proj = nn.Linear(feature_dim, hidden_dim) self.branch_norm = nn.LayerNorm(hidden_dim) self.branches = nn.ModuleList( [ nn.Sequential( nn.Conv1d(hidden_dim, hidden_dim, kernel_size=k, padding=k // 2, groups=hidden_dim), nn.GELU(), nn.Conv1d(hidden_dim, hidden_dim, kernel_size=1), ) for k in (3, 5, 9) ] ) self.gate = nn.Sequential(nn.LayerNorm(hidden_dim), nn.Linear(hidden_dim, hidden_dim), nn.Sigmoid()) self.drop = nn.Dropout(dropout) self.out_norm = nn.LayerNorm(hidden_dim) @staticmethod def _delta(x: torch.Tensor) -> torch.Tensor: out = torch.zeros_like(x) out[:, 1:] = x[:, 1:] - x[:, :-1] return out def forward(self, x: torch.Tensor) -> torch.Tensor: velocity = self._delta(x) acceleration = self._delta(velocity) h = self.pose_proj(x) + self.vel_proj(velocity) + 0.5 * self.acc_proj(acceleration) z = self.branch_norm(h).transpose(1, 2) multi = torch.stack([branch(z).transpose(1, 2) for branch in self.branches], dim=0).mean(dim=0) h = h + self.drop(multi * self.gate(h)) return self.out_norm(h) class MarginHead(nn.Module): """Linear-compatible ArcFace/CosFace head. During training, labels can be supplied to apply the margin. During evaluation, labels are omitted and the head returns scaled cosine logits. """ def __init__( self, in_features: int, num_classes: int, *, head_type: str, scale: float, margin: float, ) -> None: super().__init__() self.head_type = head_type self.scale = scale self.margin = margin self.weight = nn.Parameter(torch.empty(num_classes, in_features)) nn.init.xavier_uniform_(self.weight) def forward(self, features: torch.Tensor, labels: Optional[torch.Tensor] = None) -> torch.Tensor: cosine = F.linear(F.normalize(features), F.normalize(self.weight)) if labels is None or self.head_type == "linear": return cosine * self.scale rows = torch.arange(labels.numel(), device=labels.device) adjusted = cosine.clone() if self.head_type == "cosface": adjusted[rows, labels] -= self.margin return adjusted * self.scale if self.head_type == "arcface": clipped = cosine.clamp(-1.0 + 1e-7, 1.0 - 1e-7) theta = torch.acos(clipped) adjusted[rows, labels] = torch.cos(theta[rows, labels] + self.margin) return adjusted * self.scale raise ValueError(f"Unknown metric head type: {self.head_type}") class ResearchUnifiedAfriSignEncoder(BaseStrongUnifiedEncoder): """Exp8-v2 encoder with metric heads and optional learnable class centers.""" metric_head_type = "arcface" margin_scale = 30.0 margin = 0.25 center_weight = 0.0 def __init__(self, **kwargs: Any) -> None: super().__init__(**kwargs) hidden_dim = kwargs["hidden_dim"] task_dims: dict[str, int] = kwargs["task_dims"] self.pose_stem = MotionAwarePoseStem(kwargs["feature_dim"], hidden_dim, kwargs["dropout"]) self.heads = nn.ModuleDict( { key: MarginHead( hidden_dim, dim, head_type=self.metric_head_type, scale=self.margin_scale, margin=self.margin, ) for key, dim in task_dims.items() } ) self.class_centers = nn.ParameterDict( { key: nn.Parameter(torch.empty(dim, hidden_dim)) for key, dim in task_dims.items() } ) for centers in self.class_centers.values(): nn.init.normal_(centers, mean=0.0, std=0.02) def logits_from_features( self, task_key: str, features: torch.Tensor, labels: Optional[torch.Tensor] = None, ) -> torch.Tensor: return self.heads[task_key](features, labels) def center_loss(self, task_key: str, features: torch.Tensor, labels: torch.Tensor) -> torch.Tensor: if self.center_weight <= 0: return features.sum() * 0.0 centers = self.class_centers[task_key][labels] return F.mse_loss(F.normalize(features), F.normalize(centers)) def forward(self, batch: dict[str, torch.Tensor], task_key: str) -> torch.Tensor: features = self.encode(batch) return self.logits_from_features(task_key, features, None) def parse_wrapper_args(argv: Sequence[str]) -> tuple[argparse.Namespace, list[str]]: parser = argparse.ArgumentParser(add_help=False) parser.add_argument("--focus-mode", choices=["pose", "rgb", "multimodal"], default="multimodal") parser.add_argument("--focus-languages", nargs="+", default=["casl", "ksl", "nsi"]) parser.add_argument("--require-eval-split", action="store_true", default=True) parser.add_argument("--allow-train-only", dest="require_eval_split", action="store_false") parser.add_argument("--metric-head", choices=["arcface", "cosface", "linear"], default="arcface") parser.add_argument("--margin-scale", type=float, default=30.0) parser.add_argument("--margin", type=float, default=0.25) parser.add_argument("--center-weight", type=float, default=0.01) return parser.parse_known_args(argv) def move_batch(batch: dict[str, torch.Tensor], device: torch.device) -> dict[str, torch.Tensor]: return {k: v.to(device, non_blocking=True) for k, v in batch.items()} def cycle_loader(loader: DataLoader): while True: for batch in loader: yield batch def research_train_epoch( model: ResearchUnifiedAfriSignEncoder, tasks: Sequence[exp8.TaskSpec], loaders: dict[str, dict[str, DataLoader]], optimizer: torch.optim.Optimizer, scheduler: Optional[torch.optim.lr_scheduler.LRScheduler], device: torch.device, args: argparse.Namespace, ) -> dict[str, Any]: model.train() train_iters = {task.key: cycle_loader(loaders[task.key]["train"]) for task in tasks} schedule: list[exp8.TaskSpec] = [] for task in tasks: n = len(task.train_rows) task_batch_size = args.rgb_batch_size if task.modality == "rgb" else args.batch_size if args.samples_per_task_per_epoch > 0: steps = max(1, math.ceil(args.samples_per_task_per_epoch / max(task_batch_size, 1))) elif args.balance_tasks: steps = max(1, math.ceil(min(n, args.max_task_samples_per_epoch) / max(task_batch_size, 1))) else: steps = max(1, math.ceil(n / max(task_batch_size, 1))) schedule.extend([task] * steps) random.shuffle(schedule) loss_sum = ce_sum = con_sum = center_sum = 0.0 correct = total = 0 per_task: dict[str, dict[str, float]] = {} pbar = tqdm(schedule, desc="train", leave=False) for task in pbar: batch = move_batch(next(train_iters[task.key]), device) y = batch["y"] optimizer.zero_grad(set_to_none=True) features = model.encode(batch) logits = model.logits_from_features(task.key, features, y) ce = F.cross_entropy(logits, y, label_smoothing=args.label_smoothing) con = torch.zeros((), device=device) if args.supcon_weight > 0: con = exp8.supervised_contrastive_loss(model.contrast_features(features), y, args.temperature) center = model.center_loss(task.key, features, y) loss = ce + (args.supcon_weight * con) + (model.center_weight * center) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip) optimizer.step() if scheduler is not None: scheduler.step() bsz = y.size(0) batch_correct = int((logits.argmax(dim=1) == y).sum().item()) loss_sum += float(loss.item()) * bsz ce_sum += float(ce.item()) * bsz con_sum += float(con.item()) * bsz center_sum += float(center.item()) * bsz correct += batch_correct total += bsz slot = per_task.setdefault(task.key, {"n": 0, "correct": 0}) slot["n"] += bsz slot["correct"] += batch_correct pbar.set_postfix(loss=loss_sum / max(total, 1), acc=correct / max(total, 1), task=task.key[:16]) for key, item in per_task.items(): item["accuracy"] = item["correct"] / max(item["n"], 1) return { "loss": loss_sum / max(total, 1), "ce_loss": ce_sum / max(total, 1), "supcon_loss": con_sum / max(total, 1), "center_loss": center_sum / max(total, 1), "accuracy": correct / max(total, 1), "n": total, "per_task": per_task, } def main() -> None: wrapper_args, remaining = parse_wrapper_args(sys.argv[1:]) sys.argv = [sys.argv[0], *remaining] ResearchUnifiedAfriSignEncoder.metric_head_type = wrapper_args.metric_head ResearchUnifiedAfriSignEncoder.margin_scale = wrapper_args.margin_scale ResearchUnifiedAfriSignEncoder.margin = wrapper_args.margin ResearchUnifiedAfriSignEncoder.center_weight = wrapper_args.center_weight wanted_langs = {lang.lower() for lang in wrapper_args.focus_languages} mode = wrapper_args.focus_mode original_collect_tasks = exp8.collect_tasks def focused_collect_tasks(args: argparse.Namespace) -> list[exp8.TaskSpec]: tasks = original_collect_tasks(args) focused: list[exp8.TaskSpec] = [] for task in tasks: if task.language_code.lower() not in wanted_langs: continue if mode == "pose" and task.modality != "pose": continue if mode == "rgb" and task.modality != "rgb": continue if mode == "multimodal" and task.modality not in {"pose", "rgb"}: continue if wrapper_args.require_eval_split and not (task.val_rows or task.test_rows): continue focused.append(task) if not focused: raise SystemExit( f"No focused KCN tasks left: mode={mode}, languages={sorted(wanted_langs)}. " "Check manifests and local caches." ) return focused exp8.collect_tasks = focused_collect_tasks exp8.train_epoch = research_train_epoch exp8_v2.StrongUnifiedAfriSignEncoder = ResearchUnifiedAfriSignEncoder exp8_v2.main() if __name__ == "__main__": main()