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
+1 -6
View File
@@ -4,19 +4,14 @@ import numpy as np
import tempfile
import shutil
import torch
import pytest
import matplotlib
from torch.utils.data import Dataset
from torch.utils.data import Dataset
from khaosz.config.model_config import ModelConfig
from khaosz.data.tokenizer import BpeTokenizer
from khaosz.model.transformer import Transformer
matplotlib.use("Agg")
class RandomDataset(Dataset):
def __init__(self, length=None, max_length=64, vocab_size=1000):
self.length = length or int(np.random.randint(100, 200))
+2 -13
View File
@@ -2,7 +2,6 @@ import torch
import tempfile
import torch.distributed as dist
from pathlib import Path
from torch.optim import AdamW
from torch.optim.lr_scheduler import CosineAnnealingLR
from khaosz.data.checkpoint import Checkpoint
@@ -28,25 +27,16 @@ def test_single_process():
checkpoint = Checkpoint(
state_dict=model.state_dict(),
epoch=3,
iteration=30,
metrics={
"loss": [0.5, 0.4, 0.3, 0.2, 0.1],
"accuracy": [0.6, 0.7, 0.8, 0.85, 0.9]
}
iteration=30
)
with tempfile.TemporaryDirectory() as tmpdir:
checkpoint.save(tmpdir, save_metric_plot=True)
checkpoint.save(tmpdir)
loaded_checkpoint = Checkpoint.load(tmpdir)
assert loaded_checkpoint.epoch == 3
assert loaded_checkpoint.iteration == 30
assert loaded_checkpoint.metrics["loss"] == [0.5, 0.4, 0.3, 0.2, 0.1]
png_files = list(Path(tmpdir).glob("*.png"))
assert png_files
def simple_training():
model = torch.nn.Linear(10, 5)
optimizer = AdamW(model.parameters(), lr=1e-3)
@@ -66,7 +56,6 @@ def simple_training():
state_dict=model.state_dict(),
epoch=2,
iteration=10,
metrics={"loss": [0.3, 0.2, 0.1]}
)
rank = get_rank()