Cross validation on
(self, model, cv_data_loader, writer, info_dict, on_batch_end=True, capped_at=5, scaler=None)
| 89 | |
| 90 | @torch.inference_mode() |
| 91 | def cv(self, model, cv_data_loader, writer, info_dict, on_batch_end=True, capped_at=5, scaler=None): |
| 92 | ''' Cross validation on |
| 93 | ''' |
| 94 | logging.info('Epoch {} Step {} on_batch_end {} CV rank {}'.format(self.epoch, self.step + 1, on_batch_end, self.rank)) |
| 95 | model.eval() |
| 96 | total_num_utts, total_loss_dict = 0, {} # avoid division by 0 |
| 97 | stop = capped_at |
| 98 | for batch_idx, batch_dict in enumerate(cv_data_loader): |
| 99 | info_dict["tag"] = "CV" |
| 100 | info_dict["step"] = self.step |
| 101 | info_dict["epoch"] = self.epoch |
| 102 | info_dict["batch_idx"] = batch_idx |
| 103 | |
| 104 | num_utts = len(batch_dict["utts"]) |
| 105 | total_num_utts += num_utts |
| 106 | |
| 107 | if capped_at>0: |
| 108 | if stop <= 0: |
| 109 | continue |
| 110 | else: |
| 111 | stop -= 1 |
| 112 | |
| 113 | with autocast(device_type='cuda', enabled=scaler is not None): |
| 114 | info_dict = batch_forward(model, batch_dict, info_dict, scaler) |
| 115 | |
| 116 | for k, v in info_dict['loss_dict'].items(): |
| 117 | if k not in total_loss_dict: |
| 118 | total_loss_dict[k] = [] |
| 119 | total_loss_dict[k].append(v.item() * num_utts) |
| 120 | log_per_step(None, info_dict) |
| 121 | |
| 122 | for k, v in total_loss_dict.items(): |
| 123 | total_loss_dict[k] = sum(v) / total_num_utts |
| 124 | info_dict['loss_dict'] = total_loss_dict |
| 125 | log_per_save(writer, info_dict) |
| 126 | model_name = 'epoch_{}_whole'.format(self.epoch) if on_batch_end else 'epoch_{}_step_{}'.format(self.epoch, self.step + 1) |
| 127 | save_model(model, model_name, info_dict) |
no test coverage detected