refactor: 修改 StepMonitorCallback, 分离职责

This commit is contained in:
2026-03-04 19:45:39 +08:00
parent b53e10aac4
commit 5713b55500
5 changed files with 48 additions and 57 deletions
+2 -26
View File
@@ -1,11 +1,9 @@
import os
import json
import torch
import torch.distributed as dist
import matplotlib.pyplot as plt
from pathlib import Path
from typing import Dict, Optional, Any
from typing import Dict, Any
from khaosz.parallel.setup import get_rank
@@ -15,17 +13,14 @@ class Checkpoint:
state_dict: Dict[str, Any],
epoch: int = 0,
iteration: int = 0,
metrics: Optional[Dict[str, list]] = None,
):
self.state_dict = state_dict
self.epoch = epoch
self.iteration = iteration
self.metrics = metrics or {}
def save(
self,
save_dir: str,
save_metric_plot: bool = True,
) -> None:
save_path = Path(save_dir)
@@ -36,14 +31,10 @@ class Checkpoint:
meta = {
"epoch": self.epoch,
"iteration": self.iteration,
"metrics": self.metrics,
}
with open(save_path / "meta.json", "w") as f:
json.dump(meta, f, indent=2)
if save_metric_plot and self.metrics:
self._plot_metrics(str(save_path))
with open(save_path / f"state_dict.pt", "wb") as f:
torch.save(self.state_dict, f)
@@ -73,19 +64,4 @@ class Checkpoint:
state_dict=state_dict,
epoch=meta["epoch"],
iteration=meta["iteration"],
metrics=meta.get("metrics", {}),
)
def _plot_metrics(self, save_dir: str):
for name, values in self.metrics.items():
if not values:
continue
plt.figure(figsize=(10, 6))
plt.plot(values, label=name)
plt.xlabel("Step")
plt.ylabel("Value")
plt.title(f"Training Metric: {name}")
plt.legend()
plt.grid(True, alpha=0.3)
plt.savefig(os.path.join(save_dir, f"{name}.png"), dpi=150, bbox_inches="tight")
plt.close()
)