MCPcopy Create free account
hub / github.com/AggieSportsAnalytics/CourtCheck / predict_all

Function predict_all

backend/training/eval_tcn.py:39–68  ·  view source on GitHub ↗

Run the model on every entry. Returns (y_true, y_pred, probs, used_entries).

(entries, model, cfg, device)

Source from the content-addressed store, hash-verified

37
38
39def predict_all(entries, model, cfg, device):
40 """Run the model on every entry. Returns (y_true, y_pred, probs, used_entries)."""
41 y_true = []
42 y_pred = []
43 all_probs = []
44 used = []
45
46 model.train(False)
47 with torch.no_grad():
48 for e in entries:
49 kp_path = e["keypoints_path"]
50 label = LABEL_TO_IDX[e["mapped_label"]]
51 kp = np.load(kp_path).astype(np.float32)
52 seq = normalize_keypoints(kp, cfg.seq_len)
53 seq = temporal_derivatives(seq, orders=DERIVATIVE_ORDERS)
54 x = torch.from_numpy(seq).unsqueeze(0).to(device)
55 logits = model(x)
56 probs = F.softmax(logits, dim=-1).squeeze(0).cpu().numpy()
57 pred = int(np.argmax(probs))
58 y_true.append(label)
59 y_pred.append(pred)
60 all_probs.append(probs)
61 used.append(e)
62
63 return (
64 np.array(y_true, dtype=np.int64),
65 np.array(y_pred, dtype=np.int64),
66 np.stack(all_probs, axis=0),
67 used,
68 )
69
70
71def per_class_prf(y_true, y_pred):

Callers 1

mainFunction · 0.85

Calls 2

normalize_keypointsFunction · 0.90
temporal_derivativesFunction · 0.90

Tested by

no test coverage detected