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.
- syncnet_python/SyncNetInstance.py +210 -0
- syncnet_python/SyncNetModel.py +99 -0
- syncnet_python/__init__.py +29 -0
- syncnet_python/cli.py +127 -0
- syncnet_python/detectors/__init__.py +1 -0
- syncnet_python/detectors/s3fd/__init__.py +66 -0
- syncnet_python/detectors/s3fd/box_utils.py +233 -0
- syncnet_python/detectors/s3fd/nets.py +177 -0
- syncnet_python/run_syncnet_pipeline_on_1example.py +28 -0
- syncnet_python/run_syncnet_pipeline_on_mocha_generation_on_mocha_bench.py +157 -0
- syncnet_python/run_syncnet_pipeline_on_your_own_model_results.py +158 -0
- syncnet_python/syncnet_pipeline.py +332 -0
- syncnet_python-0.1.0.dist-info/METADATA +150 -0
- syncnet_python-0.1.0.dist-info/RECORD +18 -0
- syncnet_python-0.1.0.dist-info/WHEEL +5 -0
- syncnet_python-0.1.0.dist-info/entry_points.txt +2 -0
- syncnet_python-0.1.0.dist-info/licenses/LICENSE +21 -0
- syncnet_python-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -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()
|