syncnet-python 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,233 @@
1
+ from itertools import product as product
2
+
3
+ import numpy as np
4
+ import torch
5
+
6
+
7
+ def nms_(dets, thresh):
8
+ """
9
+ Courtesy of Ross Girshick
10
+ [https://github.com/rbgirshick/py-faster-rcnn/blob/master/lib/nms/py_cpu_nms.py]
11
+ """
12
+ x1 = dets[:, 0]
13
+ y1 = dets[:, 1]
14
+ x2 = dets[:, 2]
15
+ y2 = dets[:, 3]
16
+ scores = dets[:, 4]
17
+
18
+ areas = (x2 - x1) * (y2 - y1)
19
+ order = scores.argsort()[::-1]
20
+
21
+ keep = []
22
+ while order.size > 0:
23
+ i = order[0]
24
+ keep.append(int(i))
25
+ xx1 = np.maximum(x1[i], x1[order[1:]])
26
+ yy1 = np.maximum(y1[i], y1[order[1:]])
27
+ xx2 = np.minimum(x2[i], x2[order[1:]])
28
+ yy2 = np.minimum(y2[i], y2[order[1:]])
29
+
30
+ w = np.maximum(0.0, xx2 - xx1)
31
+ h = np.maximum(0.0, yy2 - yy1)
32
+ inter = w * h
33
+ ovr = inter / (areas[i] + areas[order[1:]] - inter)
34
+
35
+ inds = np.where(ovr <= thresh)[0]
36
+ order = order[inds + 1]
37
+
38
+ return np.array(keep).astype(int)
39
+
40
+
41
+ def decode(loc, priors, variances):
42
+ """Decode locations from predictions using priors to undo
43
+ the encoding we did for offset regression at train time.
44
+ Args:
45
+ loc (tensor): location predictions for loc layers,
46
+ Shape: [num_priors,4]
47
+ priors (tensor): Prior boxes in center-offset form.
48
+ Shape: [num_priors,4].
49
+ variances: (list[float]) Variances of priorboxes
50
+ Return:
51
+ decoded bounding box predictions
52
+ """
53
+
54
+ boxes = torch.cat(
55
+ (
56
+ priors[:, :2] + loc[:, :2] * variances[0] * priors[:, 2:],
57
+ priors[:, 2:] * torch.exp(loc[:, 2:] * variances[1]),
58
+ ),
59
+ 1,
60
+ )
61
+ boxes[:, :2] -= boxes[:, 2:] / 2
62
+ boxes[:, 2:] += boxes[:, :2]
63
+ return boxes
64
+
65
+
66
+ def nms(boxes, scores, overlap=0.5, top_k=200):
67
+ """Apply non-maximum suppression at test time to avoid detecting too many
68
+ overlapping bounding boxes for a given object.
69
+ Args:
70
+ boxes: (tensor) The location preds for the img, Shape: [num_priors,4].
71
+ scores: (tensor) The class predscores for the img, Shape:[num_priors].
72
+ overlap: (float) The overlap thresh for suppressing unnecessary boxes.
73
+ top_k: (int) The Maximum number of box preds to consider.
74
+ Return:
75
+ The indices of the kept boxes with respect to num_priors.
76
+ """
77
+
78
+ keep = scores.new(scores.size(0)).zero_().long()
79
+ if boxes.numel() == 0:
80
+ return keep, 0
81
+ x1 = boxes[:, 0]
82
+ y1 = boxes[:, 1]
83
+ x2 = boxes[:, 2]
84
+ y2 = boxes[:, 3]
85
+ area = torch.mul(x2 - x1, y2 - y1)
86
+ v, idx = scores.sort(0) # sort in ascending order
87
+ # I = I[v >= 0.01]
88
+ idx = idx[-top_k:] # indices of the top-k largest vals
89
+ xx1 = boxes.new()
90
+ yy1 = boxes.new()
91
+ xx2 = boxes.new()
92
+ yy2 = boxes.new()
93
+ w = boxes.new()
94
+ h = boxes.new()
95
+
96
+ # keep = torch.Tensor()
97
+ count = 0
98
+ while idx.numel() > 0:
99
+ i = idx[-1] # index of current largest val
100
+ # keep.append(i)
101
+ keep[count] = i
102
+ count += 1
103
+ if idx.size(0) == 1:
104
+ break
105
+ idx = idx[:-1] # remove kept element from view
106
+ # load bboxes of next highest vals
107
+ xx1.resize_(idx.size(0))
108
+ yy1.resize_(idx.size(0))
109
+ xx2.resize_(idx.size(0))
110
+ yy2.resize_(idx.size(0))
111
+
112
+ torch.index_select(x1, 0, idx, out=xx1)
113
+ torch.index_select(y1, 0, idx, out=yy1)
114
+ torch.index_select(x2, 0, idx, out=xx2)
115
+ torch.index_select(y2, 0, idx, out=yy2)
116
+ # store element-wise max with next highest score
117
+ xx1 = torch.clamp(xx1, min=x1[i])
118
+ yy1 = torch.clamp(yy1, min=y1[i])
119
+ xx2 = torch.clamp(xx2, max=x2[i])
120
+ yy2 = torch.clamp(yy2, max=y2[i])
121
+ w.resize_as_(xx2)
122
+ h.resize_as_(yy2)
123
+ w = xx2 - xx1
124
+ h = yy2 - yy1
125
+ # check sizes of xx1 and xx2.. after each iteration
126
+ w = torch.clamp(w, min=0.0)
127
+ h = torch.clamp(h, min=0.0)
128
+ inter = w * h
129
+ # IoU = i / (area(a) + area(b) - i)
130
+ rem_areas = torch.index_select(area, 0, idx) # load remaining areas)
131
+ union = (rem_areas - inter) + area[i]
132
+ IoU = inter / union # store result in iou
133
+ # keep only elements with an IoU <= overlap
134
+ idx = idx[IoU.le(overlap)]
135
+ return keep, count
136
+
137
+
138
+ class Detect(object):
139
+ def __init__(
140
+ self,
141
+ num_classes=2,
142
+ top_k=750,
143
+ nms_thresh=0.3,
144
+ conf_thresh=0.05,
145
+ variance=[0.1, 0.2],
146
+ nms_top_k=5000,
147
+ ):
148
+ self.num_classes = num_classes
149
+ self.top_k = top_k
150
+ self.nms_thresh = nms_thresh
151
+ self.conf_thresh = conf_thresh
152
+ self.variance = variance
153
+ self.nms_top_k = nms_top_k
154
+
155
+ def forward(self, loc_data, conf_data, prior_data):
156
+ num = loc_data.size(0)
157
+ num_priors = prior_data.size(0)
158
+
159
+ conf_preds = conf_data.view(num, num_priors, self.num_classes).transpose(2, 1)
160
+ batch_priors = prior_data.view(-1, num_priors, 4).expand(num, num_priors, 4)
161
+ batch_priors = batch_priors.contiguous().view(-1, 4)
162
+
163
+ decoded_boxes = decode(loc_data.view(-1, 4), batch_priors, self.variance)
164
+ decoded_boxes = decoded_boxes.view(num, num_priors, 4)
165
+
166
+ output = torch.zeros(num, self.num_classes, self.top_k, 5)
167
+
168
+ for i in range(num):
169
+ boxes = decoded_boxes[i].clone()
170
+ conf_scores = conf_preds[i].clone()
171
+
172
+ for cl in range(1, self.num_classes):
173
+ c_mask = conf_scores[cl].gt(self.conf_thresh)
174
+ scores = conf_scores[cl][c_mask]
175
+
176
+ if scores.dim() == 0:
177
+ continue
178
+ l_mask = c_mask.unsqueeze(1).expand_as(boxes)
179
+ boxes_ = boxes[l_mask].view(-1, 4)
180
+ ids, count = nms(boxes_, scores, self.nms_thresh, self.nms_top_k)
181
+ count = count if count < self.top_k else self.top_k
182
+
183
+ output[i, cl, :count] = torch.cat(
184
+ (scores[ids[:count]].unsqueeze(1), boxes_[ids[:count]]), 1
185
+ )
186
+
187
+ return output
188
+
189
+
190
+ class PriorBox(object):
191
+ def __init__(
192
+ self,
193
+ input_size,
194
+ feature_maps,
195
+ variance=[0.1, 0.2],
196
+ min_sizes=[16, 32, 64, 128, 256, 512],
197
+ steps=[4, 8, 16, 32, 64, 128],
198
+ clip=False,
199
+ ):
200
+ super(PriorBox, self).__init__()
201
+
202
+ self.imh = input_size[0]
203
+ self.imw = input_size[1]
204
+ self.feature_maps = feature_maps
205
+
206
+ self.variance = variance
207
+ self.min_sizes = min_sizes
208
+ self.steps = steps
209
+ self.clip = clip
210
+
211
+ def forward(self):
212
+ mean = []
213
+ for k, fmap in enumerate(self.feature_maps):
214
+ feath = fmap[0]
215
+ featw = fmap[1]
216
+ for i, j in product(range(feath), range(featw)):
217
+ f_kw = self.imw / self.steps[k]
218
+ f_kh = self.imh / self.steps[k]
219
+
220
+ cx = (j + 0.5) / f_kw
221
+ cy = (i + 0.5) / f_kh
222
+
223
+ s_kw = self.min_sizes[k] / self.imw
224
+ s_kh = self.min_sizes[k] / self.imh
225
+
226
+ mean += [cx, cy, s_kw, s_kh]
227
+
228
+ output = torch.FloatTensor(mean).view(-1, 4)
229
+
230
+ if self.clip:
231
+ output.clamp_(max=1, min=0)
232
+
233
+ return output
@@ -0,0 +1,177 @@
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.nn.functional as F
4
+ import torch.nn.init as init
5
+
6
+ from .box_utils import Detect, PriorBox
7
+
8
+
9
+ class L2Norm(nn.Module):
10
+ def __init__(self, n_channels, scale):
11
+ super(L2Norm, self).__init__()
12
+ self.n_channels = n_channels
13
+ self.gamma = scale or None
14
+ self.eps = 1e-10
15
+ self.weight = nn.Parameter(torch.Tensor(self.n_channels))
16
+ self.reset_parameters()
17
+
18
+ def reset_parameters(self):
19
+ init.constant_(self.weight, self.gamma)
20
+
21
+ def forward(self, x):
22
+ norm = x.pow(2).sum(dim=1, keepdim=True).sqrt() + self.eps
23
+ x = torch.div(x, norm)
24
+ out = self.weight.unsqueeze(0).unsqueeze(2).unsqueeze(3).expand_as(x) * x
25
+ return out
26
+
27
+
28
+ class S3FDNet(nn.Module):
29
+ def __init__(self, device="cuda"):
30
+ super(S3FDNet, self).__init__()
31
+ self.device = device
32
+
33
+ self.vgg = nn.ModuleList(
34
+ [
35
+ nn.Conv2d(3, 64, 3, 1, padding=1),
36
+ nn.ReLU(inplace=True),
37
+ nn.Conv2d(64, 64, 3, 1, padding=1),
38
+ nn.ReLU(inplace=True),
39
+ nn.MaxPool2d(2, 2),
40
+ nn.Conv2d(64, 128, 3, 1, padding=1),
41
+ nn.ReLU(inplace=True),
42
+ nn.Conv2d(128, 128, 3, 1, padding=1),
43
+ nn.ReLU(inplace=True),
44
+ nn.MaxPool2d(2, 2),
45
+ nn.Conv2d(128, 256, 3, 1, padding=1),
46
+ nn.ReLU(inplace=True),
47
+ nn.Conv2d(256, 256, 3, 1, padding=1),
48
+ nn.ReLU(inplace=True),
49
+ nn.Conv2d(256, 256, 3, 1, padding=1),
50
+ nn.ReLU(inplace=True),
51
+ nn.MaxPool2d(2, 2, ceil_mode=True),
52
+ nn.Conv2d(256, 512, 3, 1, padding=1),
53
+ nn.ReLU(inplace=True),
54
+ nn.Conv2d(512, 512, 3, 1, padding=1),
55
+ nn.ReLU(inplace=True),
56
+ nn.Conv2d(512, 512, 3, 1, padding=1),
57
+ nn.ReLU(inplace=True),
58
+ nn.MaxPool2d(2, 2),
59
+ nn.Conv2d(512, 512, 3, 1, padding=1),
60
+ nn.ReLU(inplace=True),
61
+ nn.Conv2d(512, 512, 3, 1, padding=1),
62
+ nn.ReLU(inplace=True),
63
+ nn.Conv2d(512, 512, 3, 1, padding=1),
64
+ nn.ReLU(inplace=True),
65
+ nn.MaxPool2d(2, 2),
66
+ nn.Conv2d(512, 1024, 3, 1, padding=6, dilation=6),
67
+ nn.ReLU(inplace=True),
68
+ nn.Conv2d(1024, 1024, 1, 1),
69
+ nn.ReLU(inplace=True),
70
+ ]
71
+ )
72
+
73
+ self.L2Norm3_3 = L2Norm(256, 10)
74
+ self.L2Norm4_3 = L2Norm(512, 8)
75
+ self.L2Norm5_3 = L2Norm(512, 5)
76
+
77
+ self.extras = nn.ModuleList(
78
+ [
79
+ nn.Conv2d(1024, 256, 1, 1),
80
+ nn.Conv2d(256, 512, 3, 2, padding=1),
81
+ nn.Conv2d(512, 128, 1, 1),
82
+ nn.Conv2d(128, 256, 3, 2, padding=1),
83
+ ]
84
+ )
85
+
86
+ self.loc = nn.ModuleList(
87
+ [
88
+ nn.Conv2d(256, 4, 3, 1, padding=1),
89
+ nn.Conv2d(512, 4, 3, 1, padding=1),
90
+ nn.Conv2d(512, 4, 3, 1, padding=1),
91
+ nn.Conv2d(1024, 4, 3, 1, padding=1),
92
+ nn.Conv2d(512, 4, 3, 1, padding=1),
93
+ nn.Conv2d(256, 4, 3, 1, padding=1),
94
+ ]
95
+ )
96
+
97
+ self.conf = nn.ModuleList(
98
+ [
99
+ nn.Conv2d(256, 4, 3, 1, padding=1),
100
+ nn.Conv2d(512, 2, 3, 1, padding=1),
101
+ nn.Conv2d(512, 2, 3, 1, padding=1),
102
+ nn.Conv2d(1024, 2, 3, 1, padding=1),
103
+ nn.Conv2d(512, 2, 3, 1, padding=1),
104
+ nn.Conv2d(256, 2, 3, 1, padding=1),
105
+ ]
106
+ )
107
+
108
+ self.softmax = nn.Softmax(dim=-1)
109
+ self.detect = Detect()
110
+
111
+ def forward(self, x):
112
+ x = x.to(self.device)
113
+ size = x.size()[2:]
114
+ sources = list()
115
+ loc = list()
116
+ conf = list()
117
+
118
+ for k in range(16):
119
+ x = self.vgg[k](x)
120
+ s = self.L2Norm3_3(x)
121
+ sources.append(s)
122
+
123
+ for k in range(16, 23):
124
+ x = self.vgg[k](x)
125
+ s = self.L2Norm4_3(x)
126
+ sources.append(s)
127
+
128
+ for k in range(23, 30):
129
+ x = self.vgg[k](x)
130
+ s = self.L2Norm5_3(x)
131
+ sources.append(s)
132
+
133
+ for k in range(30, len(self.vgg)):
134
+ x = self.vgg[k](x)
135
+ sources.append(x)
136
+
137
+ # apply extra layers and cache source layer outputs
138
+ for k, v in enumerate(self.extras):
139
+ x = F.relu(v(x), inplace=True)
140
+ if k % 2 == 1:
141
+ sources.append(x)
142
+
143
+ # apply multibox head to source layers
144
+ loc_x = self.loc[0](sources[0])
145
+ conf_x = self.conf[0](sources[0])
146
+
147
+ max_conf, _ = torch.max(conf_x[:, 0:3, :, :], dim=1, keepdim=True)
148
+ conf_x = torch.cat((max_conf, conf_x[:, 3:, :, :]), dim=1)
149
+
150
+ loc.append(loc_x.permute(0, 2, 3, 1).contiguous())
151
+ conf.append(conf_x.permute(0, 2, 3, 1).contiguous())
152
+
153
+ for i in range(1, len(sources)):
154
+ x = sources[i]
155
+ conf.append(self.conf[i](x).permute(0, 2, 3, 1).contiguous())
156
+ loc.append(self.loc[i](x).permute(0, 2, 3, 1).contiguous())
157
+
158
+ features_maps = []
159
+ for i in range(len(loc)):
160
+ feat = []
161
+ feat += [loc[i].size(1), loc[i].size(2)]
162
+ features_maps += [feat]
163
+
164
+ loc = torch.cat([o.view(o.size(0), -1) for o in loc], 1)
165
+ conf = torch.cat([o.view(o.size(0), -1) for o in conf], 1)
166
+
167
+ with torch.no_grad():
168
+ self.priorbox = PriorBox(size, features_maps)
169
+ self.priors = self.priorbox.forward()
170
+
171
+ output = self.detect.forward(
172
+ loc.view(loc.size(0), -1, 4),
173
+ self.softmax(conf.view(conf.size(0), -1, 2)),
174
+ self.priors.type(type(x.data)).to(self.device),
175
+ )
176
+
177
+ return output
@@ -0,0 +1,28 @@
1
+ from syncnet_pipeline import SyncNetPipeline # file you just saved
2
+ import logging
3
+ logging.basicConfig(
4
+ level=logging.INFO, # Change to DEBUG if you want more detail
5
+ format="%(asctime)s [%(levelname)s] %(message)s"
6
+ )
7
+
8
+ # 1. Initialise the pipeline (put your weight files in the same folder or give absolute paths)
9
+ pipe = SyncNetPipeline(
10
+ {
11
+ "s3fd_weights": "../weights/sfd_face.pth",
12
+ "syncnet_weights": "../weights/syncnet_v2.model",
13
+ },
14
+ device="cuda", # or "cpu"
15
+ )
16
+
17
+ # 2. Run inference on one clip
18
+ results = pipe.inference(
19
+ video_path="../example/video.avi", # RGB video
20
+ audio_path="../example/speech.wav", # speech track (any ffmpeg-readable format)
21
+ cache_dir="../example/cache", # optional; omit to auto-cleanup intermediates
22
+ )
23
+
24
+ # 3. Inspect outputs
25
+ offsets, confs, dists, max_conf, min_dist, s3fd_json, has_face = results
26
+ print("best-confidence :", max_conf)
27
+ print("lowest distance :", min_dist)
28
+ print("per-crop offsets :", offsets)
@@ -0,0 +1,157 @@
1
+ #!/usr/bin/env python
2
+ # -*- coding: utf-8 -*-
3
+ """
4
+ Batch SyncNet evaluation for MoChaBench.
5
+ """
6
+
7
+ import json
8
+ import csv
9
+ from pathlib import Path
10
+ import time
11
+ from collections import defaultdict
12
+
13
+ import pandas as pd
14
+ from syncnet_pipeline import SyncNetPipeline
15
+
16
+ # ------------------------------------------------------------------ #
17
+ # 1) paths & config #
18
+ # ------------------------------------------------------------------ #
19
+
20
+ # folder that contains *this* script
21
+ ROOT = Path(__file__).resolve().parents[2] # ROOT = Path(r"/full/path/to/MoChaBench") <<< fill with your own repo path/MoChaBench
22
+
23
+ # ---- adapt these if you move folders around ---------------------- #
24
+ BASE_VIDEO = ROOT / "mocha-generation" # e.g. …/MoChaBench/mocha-generation
25
+ BASE_BENCHMARK = ROOT / "benchmark" # parent folder that contains 'speeches/'
26
+ CSV_FILE = BASE_BENCHMARK / "benchmark.csv"
27
+ # ------------------------------------------------------------------ #
28
+
29
+ OUT_CSV = ROOT / "eval-lipsync" / "mocha-eval-results" / "sync_scores.csv"
30
+ OUT_JSON = ROOT / "eval-lipsync"/ "mocha-eval-results" / "sync_scores.json"
31
+
32
+ pipe = SyncNetPipeline(
33
+ {
34
+ "s3fd_weights": ROOT / "eval-lipsync" / "weights" / "sfd_face.pth",
35
+ "syncnet_weights": ROOT / "eval-lipsync" /"weights" / "syncnet_v2.model",
36
+ },
37
+ device="cuda",
38
+ )
39
+
40
+ # category buckets for your extra means
41
+ ENGLISH_1P_CATEGORIES = {
42
+ "1p_closeup_facingcamera",
43
+ "1p_camera_movement",
44
+ "1p_emotion",
45
+ "1p_mediumshot_actioncontrol",
46
+ "2p_1clip_1talk",
47
+ "1p_protrait",
48
+ }
49
+ TURNTALK_CATEGORY = {"2p_2clip_2talk"}
50
+
51
+ # ------------------------------------------------------------------ #
52
+ # 2) helper: run one sample #
53
+ # ------------------------------------------------------------------ #
54
+ def run_sample(row):
55
+ idx = int(row["idx_in_category"])
56
+ cat = row["category"].strip()
57
+ base_name = row["context_id"].strip() # e.g. 1_man_bag_of_gold
58
+ id = f"{cat}_{base_name}"
59
+ video_fp = BASE_VIDEO / cat / f"{base_name}.mp4"
60
+ audio_fp = BASE_BENCHMARK / "speeches" / cat / f"{base_name}_speech.wav"
61
+
62
+ if not video_fp.exists():
63
+ raise FileNotFoundError(f"Video not found: {video_fp}")
64
+ if not audio_fp.exists():
65
+ raise FileNotFoundError(f"Audio not found: {audio_fp}")
66
+
67
+ t0 = time.time()
68
+ off, confs, dists, best_conf, min_dist, _, has_face = pipe.inference(
69
+ video_path=str(video_fp),
70
+ audio_path=str(audio_fp),
71
+ cache_dir= ROOT / "eval-lipsync"/ "mocha-eval-results" / "cache" / id
72
+ )
73
+ return {
74
+ "idx": idx,
75
+ "category": cat,
76
+ "video": str(Path(cat) / f"{base_name}.mp4"),
77
+ "audio": str(Path(cat) / f"{base_name}_speech.wav"),
78
+ "offsets": [int(o) for o in off],
79
+ "best_conf": float(best_conf),
80
+ "min_dist": float(min_dist),
81
+ "has_face": has_face,
82
+ "runtime_s": round(time.time() - t0, 2),
83
+ }
84
+
85
+
86
+ # ------------------------------------------------------------------ #
87
+ # 3) main loop #
88
+ # ------------------------------------------------------------------ #
89
+ def main():
90
+ df = pd.read_csv(CSV_FILE)
91
+ results = []
92
+ for _, row in df.iterrows():
93
+ try:
94
+ res = run_sample(row)
95
+ results.append(res)
96
+ print(
97
+ f"[{res['idx']}] {res['category']} "
98
+ f"Δ={res['offsets']} conf={res['best_conf']:.3f} "
99
+ f"dist={res['min_dist']:.3f}"
100
+ )
101
+ except Exception as e:
102
+ print(f"[{row['idx']}] ERROR – {e}")
103
+
104
+ # ------------------------------------------------------------------ #
105
+ # 4) save full table #
106
+ # ------------------------------------------------------------------ #
107
+ pd.DataFrame(results).to_csv(OUT_CSV, index=False)
108
+ with open(OUT_JSON, "w") as f:
109
+ json.dump(results, f, indent=2)
110
+
111
+ # ------------------------------------------------------------------ #
112
+ # 5) category aggregates #
113
+ # ------------------------------------------------------------------ #
114
+ # Separate stats for each category (only when face is detected)
115
+ cat_dists = defaultdict(list)
116
+ cat_confs = defaultdict(list)
117
+
118
+ for r in results:
119
+ # !!!! sometimes SyncNetPipeline fails to detect faces, we should not inlcude those
120
+ if r["has_face"]:
121
+ cat_dists[r["category"]].append(r["min_dist"])
122
+ cat_confs[r["category"]].append(r["best_conf"])
123
+
124
+ def mean(vals):
125
+ return sum(vals) / len(vals) if vals else float("nan")
126
+
127
+ print("\n=== per-category averages (only if face detected) ===")
128
+ for cat in sorted(set(cat_dists.keys()) | set(cat_confs.keys())):
129
+ avg_dist = mean(cat_dists[cat])
130
+ avg_conf = mean(cat_confs[cat])
131
+ print(f"{cat:30s}: dist={avg_dist:.3f} conf={avg_conf:.3f} (n={len(cat_dists[cat])})")
132
+
133
+
134
+ # super-sets
135
+ english_dists = [
136
+ d for c, v in cat_dists.items() if c in ENGLISH_1P_CATEGORIES for d in v
137
+ ]
138
+ english_confs = [
139
+ c for cat, v in cat_confs.items() if cat in ENGLISH_1P_CATEGORIES for c in v
140
+ ]
141
+
142
+ dialog_dists = [
143
+ d for c, v in cat_dists.items() if c in TURNTALK_CATEGORY for d in v
144
+ ]
145
+ dialog_confs = [
146
+ c for cat, v in cat_confs.items() if cat in TURNTALK_CATEGORY for c in v
147
+ ]
148
+ print("\n--- aggregate groups (only face-detected entries) ---")
149
+ print(f"single-character English (1p*): dist={mean(english_dists):.3f} conf={mean(english_confs):.3f}")
150
+ print(f"turn-based dialogue English (2p_2clip): dist={mean(dialog_dists):.3f} conf={mean(dialog_confs):.3f}")
151
+
152
+ print(f"\nSaved detailed table → {OUT_CSV}")
153
+ print(f"Saved JSON dump → {OUT_JSON}")
154
+
155
+
156
+ if __name__ == "__main__":
157
+ main()