Report core. Args: model: Model instance or model name. data_loader: TODO. device: Target device ("cuda:0", "cpu", etc.).
(self, model, data_loader, device)
| 147 | return num |
| 148 | |
| 149 | def report_core(self, model, data_loader, device): |
| 150 | """Report core. |
| 151 | |
| 152 | Args: |
| 153 | model: Model instance or model name. |
| 154 | data_loader: TODO. |
| 155 | device: Target device ("cuda:0", "cpu", etc.). |
| 156 | """ |
| 157 | res = {} |
| 158 | for item in metrics: |
| 159 | res[item[0]] = 0.0 |
| 160 | res[item[1]] = 0.0 |
| 161 | with torch.no_grad(): |
| 162 | loss_s = 0.0 |
| 163 | uidx = 0 |
| 164 | for xs, ts, orders in data_loader: |
| 165 | xs = [x.to(device) for x in xs] |
| 166 | ts = [t.to(device) for t in ts] |
| 167 | orders = [o.to(device) for o in orders] |
| 168 | loss, pit_loss, mpit_loss, att_loss, ys, logits, labels, attractors = model( |
| 169 | xs, ts, orders |
| 170 | ) |
| 171 | loss_s += loss.item() |
| 172 | uidx += 1 |
| 173 | |
| 174 | for logit, t, att in zip(logits, labels, attractors): |
| 175 | pred = torch.argmax(torch.softmax(logit, dim=-1), dim=-1) # (T, ) |
| 176 | oov_index = torch.where(pred == self.mapping_dict["oov"])[0] |
| 177 | for i in oov_index: |
| 178 | if i > 0: |
| 179 | pred[i] = pred[i - 1] |
| 180 | else: |
| 181 | pred[i] = 0 |
| 182 | pred = [self.inv_mapping_func(i, self.mapping_dict) for i in pred] |
| 183 | decisions = [bin(num)[2:].zfill(self.max_n_speaker)[::-1] for num in pred] |
| 184 | decisions = ( |
| 185 | torch.from_numpy( |
| 186 | np.stack([np.array([int(i) for i in dec]) for dec in decisions], axis=0) |
| 187 | ) |
| 188 | .to(att.device) |
| 189 | .to(torch.float32) |
| 190 | ) |
| 191 | decisions = decisions[:, : att.shape[0]] |
| 192 | |
| 193 | stats = self.calc_diarization_error(decisions, t) |
| 194 | res["speaker_scored"] += stats["speaker_scored"] |
| 195 | res["speech_scored"] += stats["speech_scored"] |
| 196 | res["frames"] += stats["frames"] |
| 197 | for item in metrics: |
| 198 | res[item[0]] += stats[item[0]] |
| 199 | loss_s /= uidx |
| 200 | vad_acc = 0 |
| 201 | |
| 202 | return res, loss_s, stats.keys(), vad_acc |
| 203 | |
| 204 | def calc_diarization_error(self, decisions, label, label_delay=0): |
| 205 | """Calc diarization error. |
no test coverage detected