kiduki-v2-stage-gru / train_gru_stage_only.py
misakidesu's picture
Upload KiDUKi v2 stage GRU model
331b0ea verified
Raw
History Blame Contribute Delete
11.5 kB
#!/usr/bin/env python3
"""Train a stage-only GRU model on per-clip features."""
from __future__ import annotations
import argparse
import csv
import json
import random
from collections import Counter
from dataclasses import dataclass
from pathlib import Path
from typing import Dict, List, Tuple
import numpy as np
import torch
from torch import nn
from torch.utils.data import DataLoader, Dataset
ROOT = Path(__file__).resolve().parents[1]
DEFAULT_MANIFEST = ROOT / "data/annotations/feature_manifest_v2.csv"
DEFAULT_OUTPUT_DIR = ROOT / "experiments/gru_stage_only"
@dataclass
class ManifestRow:
clip_id: str
split: str
stage_label: str
feature_path: Path
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Train stage-only GRU baseline.")
parser.add_argument("--manifest", type=Path, default=DEFAULT_MANIFEST)
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR)
parser.add_argument("--epochs", type=int, default=10)
parser.add_argument("--batch-size", type=int, default=32)
parser.add_argument("--hidden-size", type=int, default=256)
parser.add_argument("--lr", type=float, default=1e-3)
parser.add_argument("--weight-decay", type=float, default=1e-4)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--device", type=str, default="cpu")
parser.add_argument("--num-workers", type=int, default=0)
parser.add_argument(
"--class-weight-mode",
type=str,
default="none",
choices=["none", "inverse", "sqrt_inverse", "effective_num"],
help=(
"Class weighting strategy for train loss. "
"'sqrt_inverse' is usually a stable default for imbalanced labels."
),
)
parser.add_argument(
"--class-weight-beta",
type=float,
default=0.999,
help="Beta for effective_num weighting (used when --class-weight-mode=effective_num).",
)
parser.add_argument("--limit-train", type=int, default=None)
parser.add_argument("--limit-val", type=int, default=None)
parser.add_argument("--limit-test", type=int, default=None)
return parser.parse_args()
def set_seed(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
def resolve_feature_path(path_str: str, manifest_path: Path) -> Path:
path = Path(path_str)
if path.is_absolute():
return path
return (manifest_path.parent / path).resolve()
def read_manifest(path: Path) -> List[ManifestRow]:
rows: List[ManifestRow] = []
with path.open("r", encoding="utf-8-sig", newline="") as fh:
reader = csv.DictReader(fh)
required = {"clip_id", "split", "stage_label", "feature_path"}
missing = required - set(reader.fieldnames or [])
if missing:
raise ValueError(f"Manifest is missing required columns: {sorted(missing)}")
for row in reader:
rows.append(
ManifestRow(
clip_id=row["clip_id"].strip(),
split=row["split"].strip(),
stage_label=row["stage_label"].strip(),
feature_path=resolve_feature_path(row["feature_path"].strip(), path),
)
)
return rows
def build_stage_map(rows: List[ManifestRow]) -> Dict[str, int]:
stage_values = sorted({row.stage_label for row in rows if row.stage_label})
if not stage_values:
raise ValueError("No stage labels found in manifest.")
return {value: idx for idx, value in enumerate(stage_values)}
class ClipFeatureDataset(Dataset):
def __init__(
self,
rows: List[ManifestRow],
stage_map: Dict[str, int],
split: str,
limit: int | None = None,
):
filtered = [row for row in rows if row.split == split]
if limit is not None:
filtered = filtered[:limit]
self.rows = filtered
self.stage_map = stage_map
def __len__(self) -> int:
return len(self.rows)
def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]:
row = self.rows[idx]
npz = np.load(row.feature_path, allow_pickle=False)
frame_features = npz["frame_features"].astype(np.float32)
raw = torch.from_numpy(frame_features)
diff = torch.zeros_like(raw)
diff[1:] = raw[1:] - raw[:-1]
features = torch.cat([raw, diff], dim=1)
stage_id = self.stage_map[row.stage_label]
return {
"features": features,
"stage_id": torch.tensor(stage_id, dtype=torch.long),
}
class GRUStageOnly(nn.Module):
def __init__(self, input_size: int, hidden_size: int, num_stages: int):
super().__init__()
self.gru = nn.GRU(
input_size=input_size,
hidden_size=hidden_size,
batch_first=True,
num_layers=1,
)
self.dropout = nn.Dropout(0.2)
self.stage_head = nn.Linear(hidden_size, num_stages)
def forward(self, features: torch.Tensor) -> Dict[str, torch.Tensor]:
output, _ = self.gru(features)
pooled = output.mean(dim=1)
pooled = self.dropout(pooled)
return {"stage_logits": self.stage_head(pooled)}
def compute_stage_counts(rows: List[ManifestRow], split: str) -> Counter[str]:
return Counter(row.stage_label for row in rows if row.split == split)
def build_class_weights(
stage_map: Dict[str, int],
counts: Counter[str],
mode: str,
beta: float,
) -> Tuple[torch.Tensor | None, Dict[str, float]]:
if mode == "none":
return None, {}
if mode == "effective_num" and not (0.0 < beta < 1.0):
raise ValueError("--class-weight-beta must be in (0, 1) for effective_num mode.")
num_classes = len(stage_map)
weights = np.ones(num_classes, dtype=np.float32)
for label, idx in stage_map.items():
count = float(counts.get(label, 0))
if count <= 0:
weights[idx] = 0.0
continue
if mode == "inverse":
weights[idx] = 1.0 / count
elif mode == "sqrt_inverse":
weights[idx] = 1.0 / np.sqrt(count)
elif mode == "effective_num":
effective_num = 1.0 - np.power(beta, count)
weights[idx] = (1.0 - beta) / max(effective_num, 1e-12)
else:
raise ValueError(f"Unsupported class-weight mode: {mode}")
positive_mask = weights > 0
if positive_mask.any():
weights[positive_mask] = weights[positive_mask] * (
positive_mask.sum() / weights[positive_mask].sum()
)
tensor = torch.tensor(weights, dtype=torch.float32)
report = {
label: float(tensor[idx].item())
for label, idx in sorted(stage_map.items(), key=lambda x: x[1])
}
return tensor, report
def compute_accuracy(logits: torch.Tensor, targets: torch.Tensor) -> float:
if targets.numel() == 0:
return 0.0
preds = logits.argmax(dim=1)
return float((preds == targets).float().mean().item())
def run_epoch(
model: GRUStageOnly,
loader: DataLoader,
optimizer: torch.optim.Optimizer | None,
device: torch.device,
criterion: nn.Module,
) -> Dict[str, float]:
is_train = optimizer is not None
model.train(is_train)
total_loss = 0.0
stage_acc_sum = 0.0
batches = 0
for batch in loader:
features = batch["features"].to(device)
stage_id = batch["stage_id"].to(device)
outputs = model(features)
loss = criterion(outputs["stage_logits"], stage_id)
if is_train:
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_loss += float(loss.item())
stage_acc_sum += compute_accuracy(outputs["stage_logits"], stage_id)
batches += 1
if batches == 0:
return {"loss": 0.0, "stage_acc": 0.0}
return {"loss": total_loss / batches, "stage_acc": stage_acc_sum / batches}
def save_json(path: Path, data: object) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8") as fh:
json.dump(data, fh, ensure_ascii=False, indent=2)
def main() -> None:
args = parse_args()
set_seed(args.seed)
rows = read_manifest(args.manifest)
stage_map = build_stage_map(rows)
save_json(
args.output_dir / "label_maps.json",
{"stage_label": stage_map, "schema_version": "v2_stage_only"},
)
train_ds = ClipFeatureDataset(rows, stage_map, "train", args.limit_train)
val_ds = ClipFeatureDataset(rows, stage_map, "val", args.limit_val)
test_ds = ClipFeatureDataset(rows, stage_map, "test", args.limit_test)
if len(train_ds) == 0:
raise ValueError("No training samples found in manifest.")
train_counts = compute_stage_counts(train_ds.rows, "train")
class_weights_cpu, class_weight_report = build_class_weights(
stage_map=stage_map,
counts=train_counts,
mode=args.class_weight_mode,
beta=args.class_weight_beta,
)
train_loader = DataLoader(
train_ds, batch_size=args.batch_size, shuffle=True, num_workers=args.num_workers
)
val_loader = DataLoader(
val_ds, batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers
)
test_loader = DataLoader(
test_ds, batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers
)
sample = train_ds[0]
input_size = sample["features"].shape[1]
model = GRUStageOnly(
input_size=input_size,
hidden_size=args.hidden_size,
num_stages=len(stage_map),
)
device = torch.device(args.device)
model.to(device)
optimizer = torch.optim.AdamW(
model.parameters(), lr=args.lr, weight_decay=args.weight_decay
)
train_criterion = nn.CrossEntropyLoss(
weight=class_weights_cpu.to(device) if class_weights_cpu is not None else None
)
eval_criterion = nn.CrossEntropyLoss()
history = []
best_val = None
for epoch in range(1, args.epochs + 1):
train_metrics = run_epoch(model, train_loader, optimizer, device, train_criterion)
val_metrics = run_epoch(model, val_loader, None, device, eval_criterion)
epoch_metrics = {
"epoch": epoch,
"train": train_metrics,
"val": val_metrics,
}
history.append(epoch_metrics)
print(json.dumps(epoch_metrics, ensure_ascii=False), flush=True)
current_val = val_metrics["stage_acc"]
if best_val is None or current_val >= best_val:
best_val = current_val
args.output_dir.mkdir(parents=True, exist_ok=True)
torch.save(model.state_dict(), args.output_dir / "best_model.pt")
test_metrics = run_epoch(model, test_loader, None, device, eval_criterion)
summary = {
"train_size": len(train_ds),
"val_size": len(val_ds),
"test_size": len(test_ds),
"class_weight_mode": args.class_weight_mode,
"class_weight_beta": args.class_weight_beta,
"class_weights": class_weight_report,
"train_stage_counts": dict(train_counts),
"history": history,
"test": test_metrics,
}
save_json(args.output_dir / "metrics.json", summary)
print(json.dumps({"final_test": test_metrics}, ensure_ascii=False))
if __name__ == "__main__":
main()