Buckets:
| """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 | |
| 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") | |
| 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.