(model_engine, eval_dataloaders, tb_writer, step, eval_gradient_accumulation_steps)
| 194 | |
| 195 | |
| 196 | def _evaluate(model_engine, eval_dataloaders, tb_writer, step, eval_gradient_accumulation_steps): |
| 197 | pbar_total = 0 |
| 198 | for eval_dataloader in eval_dataloaders.values(): |
| 199 | pbar_total += len(eval_dataloader) * len(TIMESTEP_QUANTILES_FOR_EVAL) // eval_gradient_accumulation_steps |
| 200 | if is_main_process(): |
| 201 | print('Running eval') |
| 202 | pbar = tqdm(total=pbar_total) |
| 203 | else: |
| 204 | pbar = None |
| 205 | |
| 206 | start = time.time() |
| 207 | for name, eval_dataloader in eval_dataloaders.items(): |
| 208 | losses = [] |
| 209 | for quantile in TIMESTEP_QUANTILES_FOR_EVAL: |
| 210 | loss = evaluate_single(model_engine, eval_dataloader, eval_gradient_accumulation_steps, quantile, pbar=pbar) |
| 211 | losses.append(loss) |
| 212 | if is_main_process(): |
| 213 | tb_writer.add_scalar(f'{name}/loss_quantile_{quantile:.2f}', loss, step) |
| 214 | if wandb_enable: |
| 215 | wandb.log({f'{name}/loss_quantile_{quantile:.2f}': loss, 'step': step}) |
| 216 | avg_loss = sum(losses) / len(losses) |
| 217 | if is_main_process(): |
| 218 | tb_writer.add_scalar(f'{name}/loss', avg_loss, step) |
| 219 | if wandb_enable: |
| 220 | wandb.log({f'{name}/loss': avg_loss, 'step': step}) |
| 221 | |
| 222 | duration = time.time() - start |
| 223 | if is_main_process(): |
| 224 | tb_writer.add_scalar('eval/eval_time_sec', duration, step) |
| 225 | if wandb_enable: |
| 226 | wandb.log({'eval/eval_time_sec': duration, 'step': step}) |
| 227 | pbar.close() |
| 228 | |
| 229 | |
| 230 | def evaluate(model, model_engine, eval_dataloaders, tb_writer, step, eval_gradient_accumulation_steps, disable_block_swap): |
no test coverage detected