| 408 | |
| 409 | |
| 410 | def step_csv_logger(*args: Any, cls: Type[T] = CSVLogger, **kwargs: Any) -> T: |
| 411 | logger = cls(*args, **kwargs) |
| 412 | |
| 413 | def merge_by(dicts, key): |
| 414 | from collections import defaultdict |
| 415 | |
| 416 | out = defaultdict(dict) |
| 417 | for d in dicts: |
| 418 | if key in d: |
| 419 | out[d[key]].update(d) |
| 420 | return [v for _, v in sorted(out.items())] |
| 421 | |
| 422 | def save(self) -> None: |
| 423 | """Overridden to merge CSV by the step number.""" |
| 424 | import csv |
| 425 | |
| 426 | if not self.metrics: |
| 427 | return |
| 428 | metrics = merge_by(self.metrics, "step") |
| 429 | keys = sorted({k for m in metrics for k in m}) |
| 430 | with self._fs.open(self.metrics_file_path, "w", newline="") as f: |
| 431 | writer = csv.DictWriter(f, fieldnames=keys) |
| 432 | writer.writeheader() |
| 433 | writer.writerows(metrics) |
| 434 | |
| 435 | logger.experiment.save = MethodType(save, logger.experiment) |
| 436 | |
| 437 | return logger |
| 438 | |
| 439 | |
| 440 | def chunked_cross_entropy( |