Buckets:

KaisResearch's picture
download
raw
3.92 kB
"""Evaluate a trained checkpoint on the test split (or run inference on a clip).
Evaluate on the prepared test set:
python -m src.evaluate --resume checkpoints/best.pt --data_roots data/lrs2,data/lrs3
Predict the word in a single raw video:
python -m src.evaluate --resume checkpoints/best.pt --predict path/to/clip.mp4
"""
from __future__ import annotations
import argparse
import json
import os
import sys
if __package__ in (None, ""): # allow `python src/evaluate.py`
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import numpy as np
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from src.config import Config
from src.dataset import build_datasets
from src.model import build_model
from src.transforms import ClipTransform
from src.utils import AverageMeter, accuracy, resolve_device
def load_checkpoint(path, device):
ckpt = torch.load(path, map_location=device)
label_map = ckpt["label_map"]
cfg = Config(**{k: v for k, v in ckpt["config"].items()
if k in Config.__dataclass_fields__})
cfg.num_classes = len(label_map)
model = build_model(cfg).to(device)
model.load_state_dict(ckpt["model"])
model.eval()
return model, cfg, label_map
@torch.no_grad()
def evaluate_testset(model, cfg, device):
_, _, test_ds, _ = build_datasets(cfg)
if len(test_ds) == 0:
raise SystemExit("Test set is empty — check your data roots.")
loader = DataLoader(test_ds, batch_size=cfg.batch_size, shuffle=False,
num_workers=cfg.num_workers, pin_memory=True)
criterion = nn.CrossEntropyLoss()
loss_m, top1_m, top5_m = AverageMeter(), AverageMeter(), AverageMeter()
for clips, targets in loader:
clips, targets = clips.to(device), targets.to(device)
logits = model(clips)
top1, top5 = accuracy(logits, targets, topk=(1, 5))
loss_m.update(criterion(logits, targets).item(), clips.size(0))
top1_m.update(top1, clips.size(0))
top5_m.update(top5, clips.size(0))
print(f"[test] loss {loss_m.avg:.4f} top1 {top1_m.avg:.3f} "
f"top5 {top5_m.avg:.3f} over {len(test_ds)} clips")
@torch.no_grad()
def predict_clip(model, cfg, label_map, device, video_path):
from src.preprocess import extract_mouth_clip
inv = {v: k for k, v in label_map.items()}
# Extract at the same margin the training data uses (96 -> center-crop 88).
clip = extract_mouth_clip(video_path, cfg.image_size + 8, cfg.num_frames)
tf = ClipTransform(cfg.image_size, cfg.num_frames, train=False)
x = tf(clip).unsqueeze(0).to(device) # (1, 1, T, H, W)
probs = model(x).softmax(dim=1)[0]
top = torch.topk(probs, k=min(5, len(inv)))
print(f"Predictions for {os.path.basename(video_path)}:")
for rank, (p, idx) in enumerate(zip(top.values, top.indices), 1):
print(f" {rank}. {inv[int(idx)]:<20s} {p.item():.3f}")
def main():
ap = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("--resume", required=True, help="path to checkpoint")
ap.add_argument("--predict", default="", help="raw video to classify")
ap.add_argument("--data_roots", default="",
help="comma-separated roots, overrides the checkpoint's")
ap.add_argument("--num_workers", type=int, default=4)
ap.add_argument("--device", default="auto")
args = ap.parse_args()
device = resolve_device(args.device)
model, cfg, label_map = load_checkpoint(args.resume, device)
if args.data_roots:
cfg.data_roots = args.data_roots
cfg.num_workers = args.num_workers
if args.predict:
predict_clip(model, cfg, label_map, device, args.predict)
else:
evaluate_testset(model, cfg, device)
if __name__ == "__main__":
main()

Xet Storage Details

Size:
3.92 kB
·
Xet hash:
0745480f7060df4ee78221eade04a91e62938a0dba0540b86f7894247a9c9921

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.