"""Calibrate the popcount reformulation and regenerate the RTL. Per-dim thresholds are chosen on a balanced COCO val subsample: for a person-positive dim the split maximizing F1 under `value > t`, for a person-negative dim under `value < t`. Either way the split point is the same cut, and at inference every channel uses `>` because the negative count is subtracted. Writes per_dim_thresholds.json, then calls rtl_gen so the baked constants cannot drift from the calibration that produced them. """ import argparse import json import sys from pathlib import Path import torch sys.path.insert(0, str(Path(__file__).resolve().parent)) # repo root, for `common` from common import (COCO_ROOT, D, balanced_indices, coco_split, device, # noqa: E402 f1_sweep, person_labels, pool, prf1, write_artifact) from common.pools import BALANCED_VAL # noqa: E402 import rtl_gen # noqa: E402 HERE = Path(__file__).resolve().parent CLASSIFIER = HERE / 'classifier.json' QUANT_SCALE = 8 # INT8 fixed-point scale for the layernormed feature values def main(): ap = argparse.ArgumentParser(description=__doc__) ap.add_argument('--cache', type=Path, default=COCO_ROOT / 'val_feature_cache_768' / 'val.pt') ap.add_argument('--seed', type=int, default=0) args = ap.parse_args() dev = device() c = json.loads(CLASSIFIER.read_text()) pos_dims, neg_dims = c['pos_dims'], c['neg_dims'] all_dims = pos_dims + neg_dims n_pos = len(pos_dims) print('[load] val features and person labels', flush=True) val = torch.load(args.cache, map_location='cpu', weights_only=False) coco, _ = coco_split('val2017') ids = [int(e['img_id']) for e in val] feats = torch.stack([pool(e['spatial'].float().permute(1, 2, 0).reshape(-1, D)) for e in val]).to(dev)[:, all_dims] y = person_labels(coco, ids, dev) print(f' N={feats.shape[0]} person_rate={y.float().mean():.3f}', flush=True) sel = balanced_indices(y, args.seed) X, yb = feats[sel.to(dev)], y[sel.to(dev)] print(f'[balanced] N={len(sel)} person_rate={yb.float().mean():.3f}', flush=True) per_dim = [] for local, global_dim in enumerate(all_dims): vals = X[:, local] is_pos = local < n_pos candidates = torch.quantile(vals, torch.linspace(0.05, 0.95, 19, device=dev)) best = (0.0, 0.0) for t in candidates.tolist(): m = prf1(vals > t if is_pos else vals < t, yb) if m.f1 > best[0]: best = (m.f1, t) per_dim.append({'dim_index_in_40': local, 'dim_global': int(global_dim), 'is_pos': is_pos, 'threshold': best[1], 'threshold_int8': int(round(best[1] * QUANT_SCALE)), 'per_dim_F1': best[0]}) lo = min(p['per_dim_F1'] for p in per_dim) hi = max(p['per_dim_F1'] for p in per_dim) print(f'[per-dim] calibrated, standalone F1 range {lo:.3f} - {hi:.3f}', flush=True) bits = torch.stack([X[:, p['dim_index_in_40']] > p['threshold'] for p in per_dim], 1) diff = (bits[:, :n_pos].sum(1) - bits[:, n_pos:].sum(1)).float() best_k, best_m = 0, prf1(diff > 0, yb) for t in range(-20, 21): m = prf1(diff > t, yb) if m.f1 > best_m.f1: best_k, best_m = t, m print(f'[popcount] F1={best_m.f1:.4f} P={best_m.precision:.4f} ' f'R={best_m.recall:.4f} K={best_k}', flush=True) sums = X[:, :n_pos].sum(1) - X[:, n_pos:].sum(1) add = f1_sweep(sums, yb) print(f'[additive] F1={add.f1:.4f} P={add.precision:.4f} R={add.recall:.4f} ' f't={add.threshold:.3f}', flush=True) write_artifact(HERE / 'per_dim_thresholds.json', { 'quant_scale': QUANT_SCALE, 'per_dim_thresholds': per_dim, 'popcount': {'final_threshold': int(best_k), **best_m.asdict()}, 'additive': add.asdict(), 'F1_delta_popcount_vs_additive': best_m.f1 - add.f1, }, generator='calibrate.py', classifier=CLASSIFIER, pool=BALANCED_VAL.name, split=BALANCED_VAL.split, n_images=int(len(sel)), positive_rate=round(yb.float().mean().item(), 4), selection=BALANCED_VAL.selection, seed=args.seed) for path in rtl_gen.generate(): print(f'[rtl] wrote {path}', flush=True) print('[done]', flush=True) if __name__ == '__main__': main()