afrisign-encoder-best-models / code /experiments /exp9b_kcn_research_unified_encoder.py
luciayen's picture
Upload best AfriSign Encoder checkpoints and documentation
3d02762 verified
Raw
History Blame Contribute Delete
12.3 kB
#!/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()