MCPcopy Create free account
hub / github.com/FunAudioLLM/FunMusic / cv

Method cv

inspiremusic/utils/executor.py:91–127  ·  view source on GitHub ↗

Cross validation on

(self, model, cv_data_loader, writer, info_dict, on_batch_end=True, capped_at=5, scaler=None)

Source from the content-addressed store, hash-verified

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)

Callers 1

train_one_epochMethod · 0.95

Calls 4

batch_forwardFunction · 0.90
log_per_stepFunction · 0.90
log_per_saveFunction · 0.90
save_modelFunction · 0.90

Tested by

no test coverage detected