| |
| """Experiment 8: one mixed-modality multilingual AfriSign encoder. |
| |
| This is the paper-facing "one encoder" experiment. It trains one shared |
| representation model over the ready African sign-language streams: |
| |
| - pose/landmark word data: CASL, KSL, GhSL/GSE, NSL |
| - pose/landmark sentence data: GSL Health sentence landmarks |
| - optional cached RGB word-frame data: KSL/CASL/GhSL/SASL when present |
| - optional local image data: KSLC/NSL images when present |
| |
| The checkpoint contains one shared encoder plus dataset/task-specific heads. |
| That is intentional: the representation is shared, but the label spaces are not |
| the same across KSLC images, CASL word clips, GhSL words, GSL sentences, etc. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import csv |
| import json |
| import math |
| import random |
| import sys |
| from dataclasses import dataclass |
| from pathlib import Path, PureWindowsPath |
| from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple |
|
|
| import numpy as np |
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from torch.utils.data import DataLoader, Dataset, WeightedRandomSampler |
| from tqdm.auto import tqdm |
|
|
| ROOT = Path(__file__).resolve().parents[1] |
| sys.path.insert(0, str(ROOT)) |
|
|
| from experiments import exp2_pooled_landmark_baseline as landmark_base |
| from src.data.preprocess.image import IMAGENET_MEAN, IMAGENET_STD |
| from src.data.preprocess.video import load_jpeg_frames |
|
|
|
|
| landmark_base.LANGUAGE_CANDIDATES.setdefault( |
| "casl_si", |
| ["casl/casl_si_landmarks", "casl_si_landmarks", "casl_si"], |
| ) |
|
|
|
|
| MODALITY_IDS = {"pose": 0, "rgb": 1} |
| LEVEL_IDS = {"image": 0, "word": 1, "sentence": 2} |
|
|
|
|
| def seed_everything(seed: int) -> None: |
| random.seed(seed) |
| np.random.seed(seed) |
| torch.manual_seed(seed) |
| torch.cuda.manual_seed_all(seed) |
| torch.backends.cudnn.benchmark = True |
|
|
|
|
| def read_json(path: Path) -> dict[str, Any]: |
| return json.loads(path.read_text(encoding="utf-8")) |
|
|
|
|
| def clean_key(value: str) -> str: |
| out = [] |
| for ch in value.lower(): |
| out.append(ch if ch.isalnum() else "_") |
| return "_".join(part for part in "".join(out).split("_") if part) |
|
|
|
|
| def resolve_local_path(value: str | None, *, raw_root: Path = ROOT / "data" / "raw") -> Optional[Path]: |
| """Resolve local paths written on Windows or PSC. |
| |
| Many existing manifests were created on Windows and contain absolute |
| ``C:\\...\\africansl_encoder\\...`` paths. On PSC we remap anything after the |
| repository folder name back to the current project root. |
| """ |
| if not value: |
| return None |
|
|
| raw = str(value) |
| direct = Path(raw) |
| if direct.exists(): |
| return direct |
|
|
| normalized = raw.replace("\\", "/") |
| candidates: list[Path] = [] |
|
|
| marker = "africansl_encoder/" |
| if marker in normalized: |
| candidates.append(ROOT / normalized.split(marker, 1)[1]) |
|
|
| for marker2 in ("data/raw/", "data/processed/", "data/manifests/", "checkpoints/", "results/"): |
| if marker2 in normalized: |
| candidates.append(ROOT / normalized.split(marker2, 1)[0].split("/")[-1] / normalized.split(marker2, 1)[1]) |
| candidates.append(ROOT / marker2.rstrip("/") / normalized.split(marker2, 1)[1]) |
|
|
| if not Path(normalized).is_absolute(): |
| candidates.extend([ROOT / normalized, raw_root / normalized]) |
|
|
| try: |
| win = PureWindowsPath(raw) |
| parts = list(win.parts) |
| if "africansl_encoder" in parts: |
| idx = parts.index("africansl_encoder") |
| candidates.append(ROOT.joinpath(*parts[idx + 1 :])) |
| except Exception: |
| pass |
|
|
| for candidate in candidates: |
| if candidate.exists(): |
| return candidate |
| return None |
|
|
|
|
| def fixed_sequence(value: Any, frame_count: int, feature_dim: int) -> np.ndarray: |
| return landmark_base.landmark_to_array(value, frame_count=frame_count, feature_dim=feature_dim) |
|
|
|
|
| def load_npy_landmarks(path: Path, frame_count: int, feature_dim: int) -> np.ndarray: |
| arr = np.load(path).astype(np.float32) |
| if arr.ndim > 2: |
| arr = arr.reshape(arr.shape[0], -1) |
| return fixed_sequence(arr, frame_count=frame_count, feature_dim=feature_dim) |
|
|
|
|
| def macro_f1_score(y_true: np.ndarray, y_pred: np.ndarray, num_classes: int) -> float: |
| return landmark_base.macro_f1_score(y_true, y_pred, num_classes) |
|
|
|
|
| def topk_correct(logits: torch.Tensor, y: torch.Tensor, k: int) -> int: |
| kk = min(k, logits.size(1)) |
| pred = logits.topk(kk, dim=1).indices |
| return int((pred == y.unsqueeze(1)).any(dim=1).sum().item()) |
|
|
|
|
| @dataclass |
| class TaskSpec: |
| key: str |
| name: str |
| modality: str |
| level: str |
| language_code: str |
| source_dataset: str |
| label_to_id: dict[str, int] |
| train_rows: list[Any] |
| val_rows: list[Any] |
| test_rows: list[Any] |
| row_kind: str |
| label_field: str = "label" |
| landmark_field: str = "landmarks" |
| metadata: Optional[dict[str, Any]] = None |
|
|
| @property |
| def num_classes(self) -> int: |
| return len(self.label_to_id) |
|
|
| @property |
| def counts(self) -> dict[str, int]: |
| return {"train": len(self.train_rows), "val": len(self.val_rows), "test": len(self.test_rows)} |
|
|
|
|
| class UnifiedTaskDataset(Dataset): |
| def __init__( |
| self, |
| task: TaskSpec, |
| *, |
| split: str, |
| task_idx: int, |
| lang_idx: int, |
| mean: Optional[np.ndarray], |
| std: Optional[np.ndarray], |
| frame_count: int, |
| feature_dim: int, |
| rgb_frames: int, |
| image_size: int, |
| train: bool, |
| ) -> None: |
| self.task = task |
| self.rows = task.train_rows if split == "train" else task.val_rows if split == "val" else task.test_rows |
| self.split = split |
| self.task_idx = task_idx |
| self.lang_idx = lang_idx |
| self.mean = mean |
| self.std = std |
| self.frame_count = frame_count |
| self.feature_dim = feature_dim |
| self.rgb_frames = rgb_frames |
| self.image_size = image_size |
| self.train = train |
|
|
| def __len__(self) -> int: |
| return len(self.rows) |
|
|
| def __getitem__(self, idx: int) -> dict[str, torch.Tensor]: |
| row = self.rows[idx] |
| label = self._label(row) |
| y = self.task.label_to_id[label] |
| base = { |
| "y": torch.tensor(y, dtype=torch.long), |
| "task_idx": torch.tensor(self.task_idx, dtype=torch.long), |
| "lang_idx": torch.tensor(self.lang_idx, dtype=torch.long), |
| "modality_idx": torch.tensor(MODALITY_IDS[self.task.modality], dtype=torch.long), |
| "level_idx": torch.tensor(LEVEL_IDS[self.task.level], dtype=torch.long), |
| } |
|
|
| if self.task.modality == "pose": |
| pose = self._pose(row) |
| if self.mean is not None and self.std is not None: |
| pose = (pose - self.mean) / self.std |
| if self.train: |
| pose = augment_pose(pose) |
| base["pose"] = torch.from_numpy(pose.astype(np.float32, copy=False)) |
| return base |
|
|
| rgb = self._rgb(row) |
| base["rgb"] = torch.from_numpy(rgb.astype(np.float32, copy=False)) |
| return base |
|
|
| def _label(self, row: Any) -> str: |
| if isinstance(row, dict): |
| return str(row.get(self.task.label_field)) |
| return str(row[self.task.label_field]) |
|
|
| def _pose(self, row: Any) -> np.ndarray: |
| if self.task.row_kind == "parquet": |
| return fixed_sequence(row[self.task.landmark_field], self.frame_count, self.feature_dim) |
| if self.task.row_kind == "npy": |
| path = resolve_local_path(row.get("local_path") or row.get("file_path")) |
| if path is None: |
| raise FileNotFoundError(f"Missing landmark file for {row.get('sample_id')}: {row.get('local_path')}") |
| return load_npy_landmarks(path, self.frame_count, self.feature_dim) |
| raise ValueError(f"Task {self.task.key} is not a pose task") |
|
|
| def _rgb(self, row: dict[str, Any]) -> np.ndarray: |
| if self.task.row_kind == "frames": |
| path = resolve_local_path(row.get("frames_dir")) |
| if path is None: |
| raise FileNotFoundError(f"Missing frame directory for {row.get('sample_id')}: {row.get('frames_dir')}") |
| return load_jpeg_frames(path, num_frames=self.rgb_frames, size=self.image_size) |
|
|
| if self.task.row_kind == "image": |
| path = resolve_local_path(row.get("local_path") or row.get("file_path") or row.get("file_name")) |
| if path is None: |
| raise FileNotFoundError(f"Missing image file for {row.get('sample_id')}") |
| from PIL import Image |
|
|
| with Image.open(path) as img: |
| img = img.convert("RGB").resize((self.image_size, self.image_size), Image.BILINEAR) |
| arr = np.asarray(img, dtype=np.float32) / 255.0 |
| arr = (arr - IMAGENET_MEAN) / IMAGENET_STD |
| return np.transpose(arr, (2, 0, 1))[None, ...].astype(np.float32) |
|
|
| raise ValueError(f"Task {self.task.key} is not an RGB task") |
|
|
|
|
| def augment_pose(x: np.ndarray) -> np.ndarray: |
| out = x.astype(np.float32, copy=True) |
| if random.random() < 0.75 and out.shape[0] > 8: |
| crop_ratio = random.uniform(0.85, 1.0) |
| crop_len = max(8, int(round(out.shape[0] * crop_ratio))) |
| start = random.randint(0, max(out.shape[0] - crop_len, 0)) |
| crop = out[start : start + crop_len] |
| idx = np.linspace(0, crop.shape[0] - 1, out.shape[0], dtype=np.int64) |
| out = crop[idx] |
| if random.random() < 0.80: |
| out += np.random.normal(0.0, 0.012, size=out.shape).astype(np.float32) |
| if random.random() < 0.50 and out.shape[1] % 3 == 0: |
| points = out.shape[1] // 3 |
| pts = out.reshape(out.shape[0], points, 3) |
| mask = np.random.random((out.shape[0], points, 1)) < 0.04 |
| pts[mask.repeat(3, axis=2)] = 0.0 |
| out = pts.reshape(out.shape[0], -1) |
| return out.astype(np.float32, copy=False) |
|
|
|
|
| def collate_task_batch(items: list[dict[str, torch.Tensor]]) -> dict[str, torch.Tensor]: |
| out = { |
| "y": torch.stack([b["y"] for b in items]), |
| "task_idx": torch.stack([b["task_idx"] for b in items]), |
| "lang_idx": torch.stack([b["lang_idx"] for b in items]), |
| "modality_idx": torch.stack([b["modality_idx"] for b in items]), |
| "level_idx": torch.stack([b["level_idx"] for b in items]), |
| } |
| if "pose" in items[0]: |
| out["pose"] = torch.stack([b["pose"] for b in items]) |
| if "rgb" in items[0]: |
| out["rgb"] = torch.stack([b["rgb"] for b in items]) |
| return out |
|
|
|
|
| def build_label_map(rows: Sequence[Any], field: str) -> dict[str, int]: |
| labels = sorted({str(row.get(field) if isinstance(row, dict) else row[field]) for row in rows}) |
| return {label: i for i, label in enumerate(labels)} |
|
|
|
|
| def add_landmark_parquet_tasks(args: argparse.Namespace) -> list[TaskSpec]: |
| tasks: list[TaskSpec] = [] |
| for lang in args.pose_languages: |
| requested_lang = lang |
| try: |
| bundle = landmark_base.load_language(args.data_root, lang) |
| except FileNotFoundError: |
| if lang == "casl_si": |
| print("[warn] casl_si not found; falling back to casl for this run.") |
| lang = "casl" |
| bundle = landmark_base.load_language(args.data_root, lang) |
| else: |
| print(f"[warn] landmark parquet task {requested_lang!r} not found; skipping.") |
| continue |
| level = "image" if lang in {"nsi"} else "word" |
| key = f"pose_{level}_{requested_lang if requested_lang == 'casl_si' and lang == 'casl_si' else lang}" |
| label_to_id = bundle.label_to_id |
| train_rows = bundle.train_df.to_dict("records") |
| test_rows = bundle.test_df.to_dict("records") |
| tasks.append( |
| TaskSpec( |
| key=key, |
| name=f"{lang} {level} landmark recognition", |
| modality="pose", |
| level=level, |
| language_code="casl" if lang == "casl_si" else lang, |
| source_dataset=f"{lang}_landmarks", |
| label_to_id=label_to_id, |
| train_rows=train_rows, |
| val_rows=[], |
| test_rows=test_rows, |
| row_kind="parquet", |
| label_field=bundle.label_col, |
| landmark_field=bundle.landmark_col, |
| metadata={"lang_dir": str(bundle.lang_dir)}, |
| ) |
| ) |
| return tasks |
|
|
|
|
| def add_gsl_sentence_landmark_task(args: argparse.Namespace) -> list[TaskSpec]: |
| path = args.gsl_sentence_landmark_manifest |
| if not path.is_file(): |
| return [] |
| payload = read_json(path) |
| rows = [ |
| r |
| for r in payload.get("samples", []) |
| if r.get("label") is not None and resolve_local_path(r.get("local_path")) is not None |
| ] |
| if not rows: |
| return [] |
| train_rows = [r for r in rows if r.get("split") == "train"] |
| val_rows = [r for r in rows if r.get("split") == "val"] |
| test_rows = [r for r in rows if r.get("split") == "test"] |
| label_to_id = build_label_map(train_rows, "label") |
| val_rows = [r for r in val_rows if str(r.get("label")) in label_to_id] |
| test_rows = [r for r in test_rows if str(r.get("label")) in label_to_id] |
| return [ |
| TaskSpec( |
| key="pose_sentence_gsl_health", |
| name="GSL Health sentence landmark recognition", |
| modality="pose", |
| level="sentence", |
| language_code="gse", |
| source_dataset="gsl_health_sentences_landmarks", |
| label_to_id=label_to_id, |
| train_rows=train_rows, |
| val_rows=val_rows, |
| test_rows=test_rows, |
| row_kind="npy", |
| label_field="label", |
| landmark_field="local_path", |
| metadata={"manifest": str(path)}, |
| ) |
| ] |
|
|
|
|
| def frame_rows_from_manifest(path: Path, level: str) -> list[dict[str, Any]]: |
| if not path.is_file(): |
| return [] |
| payload = read_json(path) |
| raw_rows = list(payload.get("samples", [])) + list(payload.get("clips", [])) |
| rows: list[dict[str, Any]] = [] |
| for raw_row in raw_rows: |
| row = dict(raw_row) |
| rgb = row.get("rgb") if isinstance(row.get("rgb"), dict) else {} |
| metadata = row.get("metadata") if isinstance(row.get("metadata"), dict) else {} |
| frames_dir = row.get("frames_dir") or rgb.get("frames_dir") |
| sample_id = row.get("sample_id") or row.get("clip_id") or metadata.get("sample_id") |
| label = row.get("label") |
| split = row.get("split") |
| language_code = row.get("language_code") |
| source_dataset = row.get("source_dataset") or metadata.get("source_dataset") or path.stem |
| if not frames_dir or not label or split not in {"train", "val", "test"} or not language_code: |
| continue |
| if resolve_local_path(frames_dir) is None: |
| continue |
| row.update( |
| { |
| "sample_id": str(sample_id or f"{path.stem}_{len(rows)}"), |
| "frames_dir": frames_dir, |
| "label": label, |
| "split": split, |
| "language_code": language_code, |
| "source_dataset": source_dataset, |
| "frames_extracted": True, |
| "level": row.get("level") or level, |
| } |
| ) |
| rows.append(row) |
| return rows |
|
|
|
|
| def add_frame_tasks(args: argparse.Namespace, manifests: Sequence[Path], level: str) -> list[TaskSpec]: |
| rows: list[dict[str, Any]] = [] |
| seen: set[tuple[str, str, str]] = set() |
| for path in manifests: |
| for row in frame_rows_from_manifest(path, level): |
| key = ( |
| str(row.get("language_code")), |
| str(row.get("source_dataset")), |
| str(row.get("sample_id")), |
| ) |
| if key in seen: |
| continue |
| seen.add(key) |
| rows.append(row) |
| tasks: list[TaskSpec] = [] |
| for group in sorted({(str(r.get("language_code")), str(r.get("source_dataset"))) for r in rows}): |
| lang, source_dataset = group |
| ds_rows = [r for r in rows if str(r.get("language_code")) == lang and str(r.get("source_dataset")) == source_dataset] |
| train_rows = [r for r in ds_rows if r.get("split") == "train"] |
| val_rows = [r for r in ds_rows if r.get("split") == "val"] |
| test_rows = [r for r in ds_rows if r.get("split") == "test"] |
| if len(train_rows) < args.min_task_train: |
| continue |
| label_to_id = build_label_map(train_rows, "label") |
| val_rows = [r for r in val_rows if str(r.get("label")) in label_to_id] |
| test_rows = [r for r in test_rows if str(r.get("label")) in label_to_id] |
| tasks.append( |
| TaskSpec( |
| key=f"rgb_{level}_{clean_key(source_dataset)}", |
| name=f"{source_dataset} cached RGB {level}-frame recognition", |
| modality="rgb", |
| level=level, |
| language_code=lang, |
| source_dataset=source_dataset, |
| label_to_id=label_to_id, |
| train_rows=train_rows, |
| val_rows=val_rows, |
| test_rows=test_rows, |
| row_kind="frames", |
| metadata={"manifests": [str(p) for p in manifests]}, |
| ) |
| ) |
| return tasks |
|
|
|
|
| def add_word_frame_tasks(args: argparse.Namespace) -> list[TaskSpec]: |
| if not args.include_word_frames: |
| return [] |
| return add_frame_tasks(args, args.word_frame_manifest, "word") |
|
|
|
|
| def add_sentence_frame_tasks(args: argparse.Namespace) -> list[TaskSpec]: |
| if not args.include_sentence_frames: |
| return [] |
| return add_frame_tasks(args, args.sentence_frame_manifest, "sentence") |
|
|
| def add_image_tasks(args: argparse.Namespace) -> list[TaskSpec]: |
| if not args.include_images: |
| return [] |
| rows = [] |
| seen: set[tuple[str, str, str]] = set() |
| for path in args.image_manifest: |
| if not path.is_file(): |
| continue |
| payload = read_json(path) |
| for r in payload.get("samples", []): |
| if r.get("label") is None: |
| continue |
| |
| p = resolve_local_path(r.get("local_path") or r.get("file_path") or r.get("file_name")) |
| if p is None: |
| continue |
| row = dict(r) |
| row["local_path"] = str(p) |
| key = (str(row.get("source_dataset")), str(row.get("language_code")), str(row.get("sample_id"))) |
| if key in seen: |
| continue |
| seen.add(key) |
| rows.append(row) |
| tasks: list[TaskSpec] = [] |
| for dataset_id in sorted({str(r.get("source_dataset")) for r in rows}): |
| ds_rows = [r for r in rows if r.get("source_dataset") == dataset_id] |
| train_rows = [r for r in ds_rows if r.get("split") == "train"] |
| val_rows = [r for r in ds_rows if r.get("split") == "val"] |
| test_rows = [r for r in ds_rows if r.get("split") == "test"] |
| if len(train_rows) < args.min_task_train: |
| continue |
| label_to_id = build_label_map(train_rows, "label") |
| val_rows = [r for r in val_rows if str(r.get("label")) in label_to_id] |
| test_rows = [r for r in test_rows if str(r.get("label")) in label_to_id] |
| lang = str(train_rows[0].get("language_code") or "unknown") |
| tasks.append( |
| TaskSpec( |
| key=f"rgb_image_{clean_key(dataset_id)}", |
| name=f"{dataset_id} static image recognition", |
| modality="rgb", |
| level="image", |
| language_code=lang, |
| source_dataset=dataset_id, |
| label_to_id=label_to_id, |
| train_rows=train_rows, |
| val_rows=val_rows, |
| test_rows=test_rows, |
| row_kind="image", |
| metadata={"manifests": [str(p) for p in args.image_manifest]}, |
| ) |
| ) |
| return tasks |
|
|
| def collect_tasks(args: argparse.Namespace) -> list[TaskSpec]: |
| tasks: list[TaskSpec] = [] |
| tasks.extend(add_landmark_parquet_tasks(args)) |
| if args.include_gsl_sentence_landmarks: |
| tasks.extend(add_gsl_sentence_landmark_task(args)) |
| tasks.extend(add_word_frame_tasks(args)) |
| tasks.extend(add_sentence_frame_tasks(args)) |
| tasks.extend(add_image_tasks(args)) |
|
|
| filtered = [] |
| for task in tasks: |
| if len(task.train_rows) < args.min_task_train: |
| continue |
| if task.num_classes < 2: |
| continue |
| ensure_validation_split(task, args.val_fraction, args.seed) |
| filtered.append(task) |
| return filtered |
|
|
|
|
| def row_label(row: Any, field: str) -> str: |
| return str(row.get(field) if isinstance(row, dict) else row[field]) |
|
|
|
|
| def ensure_validation_split(task: TaskSpec, val_fraction: float, seed: int) -> None: |
| """Create a label-aware validation split when a source only has train/test.""" |
| if task.val_rows or val_fraction <= 0: |
| return |
| grouped: dict[str, list[Any]] = {} |
| for row in task.train_rows: |
| grouped.setdefault(row_label(row, task.label_field), []).append(row) |
|
|
| rng = random.Random(f"{seed}:{task.key}:val") |
| new_train: list[Any] = [] |
| new_val: list[Any] = [] |
| for _label, rows in sorted(grouped.items()): |
| rows = list(rows) |
| rng.shuffle(rows) |
| if len(rows) < 2: |
| new_train.extend(rows) |
| continue |
| n_val = max(1, int(round(len(rows) * val_fraction))) |
| n_val = min(n_val, len(rows) - 1) |
| new_val.extend(rows[:n_val]) |
| new_train.extend(rows[n_val:]) |
|
|
| if new_val: |
| task.train_rows = new_train |
| task.val_rows = new_val |
| meta = task.metadata or {} |
| meta["generated_val_from_train"] = True |
| meta["generated_val_fraction"] = val_fraction |
| task.metadata = meta |
|
|
|
|
| def compute_pose_stats(tasks: Sequence[TaskSpec], args: argparse.Namespace) -> tuple[np.ndarray, np.ndarray]: |
| total = 0 |
| sum_x = np.zeros(args.feature_dim, dtype=np.float64) |
| sum_x2 = np.zeros(args.feature_dim, dtype=np.float64) |
| rng = random.Random(args.seed) |
| for task in tasks: |
| if task.modality != "pose": |
| continue |
| rows = list(task.train_rows) |
| if args.stats_sample > 0 and len(rows) > args.stats_sample: |
| rows = rng.sample(rows, args.stats_sample) |
| for row in tqdm(rows, desc=f"pose stats {task.key}", leave=False): |
| if task.row_kind == "parquet": |
| arr = fixed_sequence(row[task.landmark_field], args.frame_count, args.feature_dim) |
| else: |
| p = resolve_local_path(row.get("local_path") or row.get("file_path")) |
| if p is None: |
| continue |
| arr = load_npy_landmarks(p, args.frame_count, args.feature_dim) |
| sum_x += arr.sum(axis=0) |
| sum_x2 += np.square(arr, dtype=np.float64).sum(axis=0) |
| total += arr.shape[0] |
| mean = sum_x / max(total, 1) |
| var = (sum_x2 / max(total, 1)) - np.square(mean) |
| std = np.sqrt(np.maximum(var, 1e-12)) |
| std = np.where(std < 1e-6, 1.0, std) |
| return mean.astype(np.float32), std.astype(np.float32) |
|
|
|
|
| class SmallFrameCNN(nn.Module): |
| def __init__(self, hidden_dim: int) -> None: |
| super().__init__() |
| self.net = nn.Sequential( |
| nn.Conv2d(3, 32, 3, stride=2, padding=1), |
| nn.BatchNorm2d(32), |
| nn.GELU(), |
| nn.Conv2d(32, 64, 3, stride=2, padding=1), |
| nn.BatchNorm2d(64), |
| nn.GELU(), |
| nn.Conv2d(64, 128, 3, stride=2, padding=1), |
| nn.BatchNorm2d(128), |
| nn.GELU(), |
| nn.Conv2d(128, hidden_dim, 3, stride=2, padding=1), |
| nn.BatchNorm2d(hidden_dim), |
| nn.GELU(), |
| nn.AdaptiveAvgPool2d(1), |
| ) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| return self.net(x).flatten(1) |
|
|
|
|
| class EfficientNetB0FrameEncoder(nn.Module): |
| """Shared ImageNet frame encoder for both RGB videos and static images.""" |
|
|
| def __init__(self, hidden_dim: int, train_backbone: str, pretrained: bool) -> None: |
| super().__init__() |
| try: |
| from torchvision.models import EfficientNet_B0_Weights, efficientnet_b0 |
| except ImportError as exc: |
| raise ImportError("torchvision is required for --rgb-backbone efficientnet_b0") from exc |
|
|
| weights = EfficientNet_B0_Weights.IMAGENET1K_V1 if pretrained else None |
| model = efficientnet_b0(weights=weights) |
| feat_dim = model.classifier[1].in_features |
| model.classifier = nn.Identity() |
|
|
| if train_backbone == "none": |
| for p in model.parameters(): |
| p.requires_grad = False |
| elif train_backbone == "last": |
| for p in model.parameters(): |
| p.requires_grad = False |
| for block in list(model.features.children())[-2:]: |
| for p in block.parameters(): |
| p.requires_grad = True |
| elif train_backbone == "all": |
| for p in model.parameters(): |
| p.requires_grad = True |
| else: |
| raise ValueError(f"Unknown train_backbone={train_backbone!r}") |
|
|
| self.backbone = model |
| self.proj = nn.Linear(feat_dim, hidden_dim) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| return self.proj(self.backbone(x)) |
|
|
|
|
| class UnifiedAfriSignEncoder(nn.Module): |
| def __init__( |
| self, |
| *, |
| task_dims: dict[str, int], |
| num_languages: int, |
| hidden_dim: int, |
| feature_dim: int, |
| max_tokens: int, |
| layers: int, |
| heads: int, |
| ff_dim: int, |
| dropout: float, |
| rgb_backbone: str, |
| rgb_train_backbone: str, |
| rgb_pretrained: bool, |
| ) -> None: |
| super().__init__() |
| self.task_keys = list(task_dims) |
| self.pose_proj = nn.Linear(feature_dim, hidden_dim) |
| if rgb_backbone == "small_cnn": |
| self.rgb_cnn = SmallFrameCNN(hidden_dim) |
| elif rgb_backbone == "efficientnet_b0": |
| self.rgb_cnn = EfficientNetB0FrameEncoder(hidden_dim, rgb_train_backbone, rgb_pretrained) |
| else: |
| raise ValueError(f"Unknown rgb_backbone={rgb_backbone!r}") |
| self.cls = nn.Parameter(torch.zeros(1, 1, hidden_dim)) |
| self.pos = nn.Embedding(max_tokens + 1, hidden_dim) |
| self.lang_emb = nn.Embedding(num_languages, hidden_dim) |
| self.modality_emb = nn.Embedding(len(MODALITY_IDS), hidden_dim) |
| self.level_emb = nn.Embedding(len(LEVEL_IDS), hidden_dim) |
| self.task_emb = nn.Embedding(len(self.task_keys), hidden_dim) |
| block = nn.TransformerEncoderLayer( |
| d_model=hidden_dim, |
| nhead=heads, |
| dim_feedforward=ff_dim, |
| dropout=dropout, |
| batch_first=True, |
| norm_first=True, |
| ) |
| self.encoder = nn.TransformerEncoder(block, layers, enable_nested_tensor=False) |
| self.norm = nn.LayerNorm(hidden_dim) |
| self.drop = nn.Dropout(dropout) |
| self.heads = nn.ModuleDict({key: nn.Linear(hidden_dim, dim) for key, dim in task_dims.items()}) |
| self.projector = nn.Sequential(nn.Linear(hidden_dim, hidden_dim), nn.GELU(), nn.Linear(hidden_dim, hidden_dim)) |
| nn.init.trunc_normal_(self.cls, std=0.02) |
|
|
| def encode(self, batch: dict[str, torch.Tensor]) -> torch.Tensor: |
| if "pose" in batch: |
| x = self.pose_proj(batch["pose"]) |
| elif "rgb" in batch: |
| rgb = batch["rgb"] |
| b, t, c, h, w = rgb.shape |
| x = self.rgb_cnn(rgb.reshape(b * t, c, h, w)).reshape(b, t, -1) |
| else: |
| raise ValueError("Batch must contain pose or rgb") |
|
|
| b, t, _ = x.shape |
| cls = self.cls.expand(b, -1, -1) |
| x = torch.cat([cls, x], dim=1) |
| positions = torch.arange(t + 1, device=x.device) |
| context = ( |
| self.lang_emb(batch["lang_idx"]) |
| + self.modality_emb(batch["modality_idx"]) |
| + self.level_emb(batch["level_idx"]) |
| + self.task_emb(batch["task_idx"]) |
| ).unsqueeze(1) |
| x = x + self.pos(positions).unsqueeze(0) + context |
| x = self.encoder(x)[:, 0] |
| return self.drop(self.norm(x)) |
|
|
| def forward(self, batch: dict[str, torch.Tensor], task_key: str) -> torch.Tensor: |
| features = self.encode(batch) |
| return self.heads[task_key](features) |
|
|
| def contrast_features(self, features: torch.Tensor) -> torch.Tensor: |
| return F.normalize(self.projector(features), dim=1) |
|
|
|
|
| def supervised_contrastive_loss(features: torch.Tensor, y: torch.Tensor, temperature: float) -> torch.Tensor: |
| if features.size(0) <= 1: |
| return features.sum() * 0.0 |
| same = y.unsqueeze(0) == y.unsqueeze(1) |
| eye = torch.eye(features.size(0), dtype=torch.bool, device=features.device) |
| positive = same & ~eye |
| anchors = positive.sum(dim=1) > 0 |
| if not torch.any(anchors): |
| return features.sum() * 0.0 |
| logits = torch.matmul(features, features.T) / max(temperature, 1e-6) |
| logits = logits - logits.max(dim=1, keepdim=True).values.detach() |
| exp_logits = torch.exp(logits) * (~eye).float() |
| log_prob = logits - torch.log(exp_logits.sum(dim=1, keepdim=True).clamp_min(1e-12)) |
| pos_log_prob = (positive.float() * log_prob).sum(dim=1) / positive.sum(dim=1).clamp_min(1) |
| return -pos_log_prob[anchors].mean() |
|
|
|
|
| def make_task_loaders( |
| tasks: Sequence[TaskSpec], |
| task_to_idx: dict[str, int], |
| lang_to_idx: dict[str, int], |
| mean: np.ndarray, |
| std: np.ndarray, |
| args: argparse.Namespace, |
| ) -> dict[str, dict[str, DataLoader]]: |
| loaders: dict[str, dict[str, DataLoader]] = {} |
| for task in tasks: |
| loaders[task.key] = {} |
| for split in ("train", "val", "test"): |
| rows = task.train_rows if split == "train" else task.val_rows if split == "val" else task.test_rows |
| if not rows: |
| continue |
| ds = UnifiedTaskDataset( |
| task, |
| split=split, |
| task_idx=task_to_idx[task.key], |
| lang_idx=lang_to_idx[task.language_code], |
| mean=mean if task.modality == "pose" else None, |
| std=std if task.modality == "pose" else None, |
| frame_count=args.frame_count, |
| feature_dim=args.feature_dim, |
| rgb_frames=args.rgb_frames, |
| image_size=args.image_size, |
| train=split == "train", |
| ) |
| sampler = None |
| if split == "train" and args.balance_classes: |
| labels = [task.label_to_id[str(row.get(task.label_field) if isinstance(row, dict) else row[task.label_field])] for row in rows] |
| counts: dict[int, int] = {} |
| for label in labels: |
| counts[label] = counts.get(label, 0) + 1 |
| weights = [1.0 / counts[label] for label in labels] |
| sampler = WeightedRandomSampler(torch.as_tensor(weights, dtype=torch.double), len(weights), replacement=True) |
| batch_size = args.rgb_batch_size if task.modality == "rgb" else args.batch_size |
| loaders[task.key][split] = DataLoader( |
| ds, |
| batch_size=batch_size, |
| shuffle=(split == "train" and sampler is None), |
| sampler=sampler, |
| num_workers=args.num_workers, |
| pin_memory=torch.cuda.is_available(), |
| collate_fn=collate_task_batch, |
| ) |
| return loaders |
|
|
|
|
| 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) -> Iterable[dict[str, torch.Tensor]]: |
| while True: |
| for batch in loader: |
| yield batch |
|
|
|
|
| def train_epoch( |
| model: UnifiedAfriSignEncoder, |
| tasks: Sequence[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[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 = 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.heads[task.key](features) |
| ce = F.cross_entropy(logits, y, label_smoothing=args.label_smoothing) |
| con = torch.zeros((), device=device) |
| if args.supcon_weight > 0: |
| con = supervised_contrastive_loss(model.contrast_features(features), y, args.temperature) |
| loss = ce + (args.supcon_weight * con) |
| 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 |
| 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), |
| "accuracy": correct / max(total, 1), |
| "n": total, |
| "per_task": per_task, |
| } |
|
|
|
|
| @torch.no_grad() |
| def evaluate_split( |
| model: UnifiedAfriSignEncoder, |
| tasks: Sequence[TaskSpec], |
| loaders: dict[str, dict[str, DataLoader]], |
| split: str, |
| device: torch.device, |
| ) -> dict[str, Any]: |
| model.eval() |
| out: dict[str, Any] = {} |
| for task in tasks: |
| loader = loaders.get(task.key, {}).get(split) |
| if loader is None: |
| continue |
| loss_sum = 0.0 |
| correct = 0 |
| top5 = 0 |
| total = 0 |
| y_true: list[int] = [] |
| y_pred: list[int] = [] |
| for batch in loader: |
| batch = move_batch(batch, device) |
| y = batch["y"] |
| logits = model(batch, task.key) |
| loss_sum += float(F.cross_entropy(logits, y).item()) * y.size(0) |
| pred = logits.argmax(dim=1) |
| correct += int((pred == y).sum().item()) |
| top5 += topk_correct(logits, y, 5) |
| total += y.size(0) |
| y_true.extend(y.detach().cpu().numpy().tolist()) |
| y_pred.extend(pred.detach().cpu().numpy().tolist()) |
| out[task.key] = { |
| "loss": loss_sum / max(total, 1), |
| "top1": correct / max(total, 1), |
| "top5": top5 / max(total, 1), |
| "macro_f1": macro_f1_score(np.asarray(y_true), np.asarray(y_pred), task.num_classes), |
| "n": total, |
| "num_classes": task.num_classes, |
| "language_code": task.language_code, |
| "source_dataset": task.source_dataset, |
| "level": task.level, |
| "modality": task.modality, |
| } |
|
|
| for group_key, predicate in { |
| "macro_all": lambda _task: True, |
| "macro_word": lambda task: task.level == "word", |
| "macro_sentence": lambda task: task.level == "sentence", |
| "macro_image": lambda task: task.level == "image", |
| "macro_pose": lambda task: task.modality == "pose", |
| "macro_rgb": lambda task: task.modality == "rgb", |
| }.items(): |
| vals = [out[t.key] for t in tasks if t.key in out and predicate(t)] |
| if vals: |
| out[group_key] = { |
| "top1": float(np.mean([v["top1"] for v in vals])), |
| "top5": float(np.mean([v["top5"] for v in vals])), |
| "macro_f1": float(np.mean([v["macro_f1"] for v in vals])), |
| "n_tasks": len(vals), |
| } |
| return out |
|
|
|
|
| def write_manifest_summary(tasks: Sequence[TaskSpec], path: Path) -> None: |
| rows = [] |
| for task in tasks: |
| rows.append( |
| { |
| "task_key": task.key, |
| "name": task.name, |
| "modality": task.modality, |
| "level": task.level, |
| "language_code": task.language_code, |
| "source_dataset": task.source_dataset, |
| "classes": task.num_classes, |
| "train": len(task.train_rows), |
| "val": len(task.val_rows), |
| "test": len(task.test_rows), |
| } |
| ) |
| if not rows: |
| return |
| with path.open("w", newline="", encoding="utf-8") as handle: |
| writer = csv.DictWriter(handle, fieldnames=list(rows[0].keys())) |
| writer.writeheader() |
| writer.writerows(rows) |
|
|
|
|
| def parse_args() -> argparse.Namespace: |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("--data-root", type=Path, default=ROOT / "data" / "raw") |
| parser.add_argument("--pose-languages", nargs="+", default=["casl_si", "ksl", "gse", "nsi"]) |
| parser.add_argument("--gsl-sentence-landmark-manifest", type=Path, default=ROOT / "data/manifests/gsl_sentence_landmarks.json") |
| parser.add_argument("--word-frame-manifest", type=Path, nargs="+", default=[ROOT / "data/manifests/word_video_frames.json"]) |
| parser.add_argument("--sentence-frame-manifest", type=Path, nargs="+", default=[ROOT / "data/manifests/sentence_video_frames.json"]) |
| parser.add_argument("--image-manifest", type=Path, nargs="+", default=[ROOT / "data/manifests/unified_images_split.json"]) |
| parser.add_argument("--out-dir", type=Path, default=ROOT / "results/exp8_unified_mixed_encoder") |
| parser.add_argument("--checkpoint-dir", type=Path, default=ROOT / "checkpoints/exp8_unified_mixed_encoder") |
| parser.add_argument("--run-name", default="unified_afrisign_encoder") |
| parser.add_argument("--epochs", type=int, default=40) |
| parser.add_argument("--batch-size", type=int, default=64) |
| parser.add_argument("--rgb-batch-size", type=int, default=8) |
| parser.add_argument("--lr", type=float, default=2e-4) |
| parser.add_argument("--weight-decay", type=float, default=1e-4) |
| parser.add_argument("--patience", type=int, default=10) |
| parser.add_argument("--seed", type=int, default=42) |
| parser.add_argument("--frame-count", type=int, default=64) |
| parser.add_argument("--feature-dim", type=int, default=225) |
| parser.add_argument("--rgb-frames", type=int, default=16) |
| parser.add_argument("--image-size", type=int, default=112) |
| parser.add_argument("--hidden-dim", type=int, default=256) |
| parser.add_argument("--layers", type=int, default=4) |
| parser.add_argument("--heads", type=int, default=8) |
| parser.add_argument("--ff-dim", type=int, default=1024) |
| parser.add_argument("--dropout", type=float, default=0.15) |
| parser.add_argument("--label-smoothing", type=float, default=0.05) |
| parser.add_argument("--supcon-weight", type=float, default=0.03) |
| parser.add_argument("--temperature", type=float, default=0.10) |
| parser.add_argument("--grad-clip", type=float, default=2.0) |
| parser.add_argument("--stats-sample", type=int, default=1000) |
| parser.add_argument("--num-workers", type=int, default=0 if sys.platform == "win32" else 2) |
| parser.add_argument("--min-task-train", type=int, default=8) |
| parser.add_argument("--val-fraction", type=float, default=0.10, help="Label-aware validation fraction when a task only has train/test splits.") |
| parser.add_argument("--samples-per-task-per-epoch", type=int, default=0) |
| parser.add_argument("--max-task-samples-per-epoch", type=int, default=4096) |
| parser.add_argument("--balance-tasks", action="store_true", default=True) |
| parser.add_argument("--no-balance-tasks", dest="balance_tasks", action="store_false") |
| parser.add_argument("--balance-classes", action="store_true", default=True) |
| parser.add_argument("--no-balance-classes", dest="balance_classes", action="store_false") |
| parser.add_argument("--include-gsl-sentence-landmarks", action="store_true", default=True) |
| parser.add_argument("--no-gsl-sentence-landmarks", dest="include_gsl_sentence_landmarks", action="store_false") |
| parser.add_argument("--include-word-frames", action="store_true", help="Use cached RGB word frames if present.") |
| parser.add_argument("--include-sentence-frames", action="store_true", help="Use cached RGB sentence frames if present.") |
| parser.add_argument("--include-images", action="store_true", help="Use local cached RGB images if present.") |
| parser.add_argument("--rgb-backbone", choices=["small_cnn", "efficientnet_b0"], default="small_cnn") |
| parser.add_argument("--rgb-train-backbone", choices=["none", "last", "all"], default="none") |
| parser.add_argument("--rgb-pretrained", action="store_true", default=True) |
| parser.add_argument("--no-rgb-pretrained", dest="rgb_pretrained", action="store_false") |
| parser.add_argument("--monitor", default="macro_all", help="Validation aggregate key to select checkpoints.") |
| parser.add_argument("--dry-run", action="store_true") |
| return parser.parse_args() |
|
|
|
|
| def main() -> None: |
| args = parse_args() |
| seed_everything(args.seed) |
| args.out_dir.mkdir(parents=True, exist_ok=True) |
| args.checkpoint_dir.mkdir(parents=True, exist_ok=True) |
|
|
| tasks = collect_tasks(args) |
| if not tasks: |
| raise SystemExit("No usable tasks found. Check local landmarks/frame caches first.") |
|
|
| language_codes = sorted({task.language_code for task in tasks}) |
| lang_to_idx = {lang: i for i, lang in enumerate(language_codes)} |
| task_to_idx = {task.key: i for i, task in enumerate(tasks)} |
|
|
| print("\nUnified AfriSign encoder setup") |
| print("run:", args.run_name) |
| print("data root:", args.data_root) |
| print("tasks:", len(tasks)) |
| print("languages:", ", ".join(language_codes)) |
| print() |
| for task in tasks: |
| print( |
| f"{task.key:32s} | {task.modality:4s} | {task.level:8s} | {task.language_code:7s} | " |
| f"classes={task.num_classes:5d} | train={len(task.train_rows):5d} " |
| f"val={len(task.val_rows):5d} test={len(task.test_rows):5d}" |
| ) |
|
|
| summary_csv = args.out_dir / f"{args.run_name}_task_summary.csv" |
| write_manifest_summary(tasks, summary_csv) |
|
|
| if args.dry_run: |
| print("\nDry run complete. No model was trained.") |
| print("Task summary:", summary_csv) |
| return |
|
|
| print("\nComputing shared pose normalization...") |
| mean, std = compute_pose_stats(tasks, args) |
| loaders = make_task_loaders(tasks, task_to_idx, lang_to_idx, mean, std, args) |
| task_dims = {task.key: task.num_classes for task in tasks} |
| max_tokens = max(args.frame_count, args.rgb_frames) |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| model = UnifiedAfriSignEncoder( |
| task_dims=task_dims, |
| num_languages=len(language_codes), |
| hidden_dim=args.hidden_dim, |
| feature_dim=args.feature_dim, |
| max_tokens=max_tokens, |
| layers=args.layers, |
| heads=args.heads, |
| ff_dim=args.ff_dim, |
| dropout=args.dropout, |
| rgb_backbone=args.rgb_backbone, |
| rgb_train_backbone=args.rgb_train_backbone, |
| rgb_pretrained=args.rgb_pretrained, |
| ).to(device) |
| params = sum(p.numel() for p in model.parameters() if p.requires_grad) |
|
|
| optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay) |
| steps_per_epoch = sum( |
| max( |
| 1, |
| math.ceil( |
| (args.samples_per_task_per_epoch or min(len(t.train_rows), args.max_task_samples_per_epoch)) |
| / max(args.rgb_batch_size if t.modality == "rgb" else args.batch_size, 1) |
| ), |
| ) |
| for t in tasks |
| ) |
| scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer, max_lr=args.lr, epochs=args.epochs, steps_per_epoch=max(steps_per_epoch, 1)) |
|
|
| print("\ndevice:", device) |
| print("trainable params:", params) |
| print("steps per epoch:", steps_per_epoch) |
|
|
| best_score = -1.0 |
| best_epoch = 0 |
| wait = 0 |
| history: list[dict[str, Any]] = [] |
| ckpt_path = args.checkpoint_dir / f"{args.run_name}_seed{args.seed}_best.pt" |
|
|
| for epoch in range(1, args.epochs + 1): |
| train_metrics = train_epoch(model, tasks, loaders, optimizer, scheduler, device, args) |
| val_metrics = evaluate_split(model, tasks, loaders, "val", device) |
| monitor_item = val_metrics.get(args.monitor) or val_metrics.get("macro_all") or {} |
| score = float(monitor_item.get("macro_f1", monitor_item.get("top1", 0.0))) |
| history.append({"epoch": epoch, "train": train_metrics, "val": val_metrics, "lr": scheduler.get_last_lr()[0]}) |
| print( |
| f"epoch {epoch:03d} train_acc={train_metrics['accuracy']:.3f} " |
| f"train_loss={train_metrics['loss']:.4f} val_{args.monitor}_f1={score:.3f}" |
| ) |
| if score > best_score: |
| best_score = score |
| best_epoch = epoch |
| wait = 0 |
| torch.save( |
| { |
| "epoch": epoch, |
| "monitor": args.monitor, |
| "monitor_score": best_score, |
| "model": model.state_dict(), |
| "mean": mean, |
| "std": std, |
| "tasks": [dict(task.__dict__, train_rows=[], val_rows=[], test_rows=[]) for task in tasks], |
| "task_to_idx": task_to_idx, |
| "lang_to_idx": lang_to_idx, |
| "args": {k: str(v) if isinstance(v, Path) else v for k, v in vars(args).items()}, |
| }, |
| ckpt_path, |
| ) |
| print(" saved best checkpoint:", ckpt_path) |
| else: |
| wait += 1 |
| if wait >= args.patience: |
| print(f"Early stopping after {args.patience} epochs without improvement.") |
| break |
|
|
| if ckpt_path.exists(): |
| state = torch.load(ckpt_path, map_location=device, weights_only=False) |
| model.load_state_dict(state["model"]) |
| final_val = evaluate_split(model, tasks, loaders, "val", device) |
| final_test = evaluate_split(model, tasks, loaders, "test", device) |
|
|
| result = { |
| "experiment": "exp8_unified_mixed_encoder", |
| "description": "One shared mixed-modality multilingual AfriSign encoder with dataset/task-specific heads.", |
| "run_name": args.run_name, |
| "seed": args.seed, |
| "best_epoch": best_epoch, |
| "monitor": args.monitor, |
| "best_monitor_score": best_score, |
| "params": params, |
| "languages": language_codes, |
| "task_summary": [ |
| { |
| "task_key": task.key, |
| "name": task.name, |
| "modality": task.modality, |
| "level": task.level, |
| "language_code": task.language_code, |
| "source_dataset": task.source_dataset, |
| "num_classes": task.num_classes, |
| "counts": task.counts, |
| } |
| for task in tasks |
| ], |
| "val": final_val, |
| "test": final_test, |
| "checkpoint": str(ckpt_path), |
| "args": {k: str(v) if isinstance(v, Path) else v for k, v in vars(args).items()}, |
| } |
| result_path = args.out_dir / f"{args.run_name}_seed{args.seed}_results.json" |
| history_path = args.out_dir / f"{args.run_name}_seed{args.seed}_history.json" |
| result_path.write_text(json.dumps(result, indent=2, default=str), encoding="utf-8") |
| history_path.write_text(json.dumps(history, indent=2, default=str), encoding="utf-8") |
|
|
| print("\nSaved:") |
| print(" result :", result_path) |
| print(" history:", history_path) |
| print(" summary:", summary_csv) |
| print(" ckpt :", ckpt_path) |
| print("\nTest macro:") |
| print(json.dumps({k: v for k, v in final_test.items() if k.startswith("macro_")}, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|
|
|
|
|
|
|