fix: 修复 dataset 和 checkpoint 的 bug
This commit is contained in:
@@ -1,12 +1,12 @@
|
||||
import os
|
||||
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
|
||||
from khaosz.parallel.setup import spawn_parallel_fn
|
||||
from khaosz.parallel.setup import get_rank, spawn_parallel_fn
|
||||
|
||||
def test_single_process():
|
||||
model = torch.nn.Linear(10, 5)
|
||||
@@ -26,8 +26,7 @@ def test_single_process():
|
||||
scheduler.step()
|
||||
|
||||
checkpoint = Checkpoint(
|
||||
optimizer_state_dict=optimizer.state_dict(),
|
||||
scheduler_state_dict=scheduler.state_dict(),
|
||||
state_dict=model.state_dict(),
|
||||
epoch=3,
|
||||
iteration=30,
|
||||
metrics={
|
||||
@@ -45,21 +44,14 @@ def test_single_process():
|
||||
assert loaded_checkpoint.iteration == 30
|
||||
assert loaded_checkpoint.metrics["loss"] == [0.5, 0.4, 0.3, 0.2, 0.1]
|
||||
|
||||
assert 'param_groups' in loaded_checkpoint.optimizer_state_dict
|
||||
assert 'state' in loaded_checkpoint.optimizer_state_dict
|
||||
|
||||
png_files = list(Path(tmpdir).glob("*.png"))
|
||||
assert png_files
|
||||
|
||||
def simple_training():
|
||||
rank = int(os.environ.get('LOCAL_RANK', 0))
|
||||
|
||||
# 简单的训练逻辑
|
||||
model = torch.nn.Linear(10, 5)
|
||||
optimizer = AdamW(model.parameters(), lr=1e-3)
|
||||
scheduler = CosineAnnealingLR(optimizer, T_max=10)
|
||||
|
||||
# 训练步骤
|
||||
for epoch in range(2):
|
||||
for iteration in range(5):
|
||||
x = torch.randn(16, 10)
|
||||
@@ -71,18 +63,29 @@ def simple_training():
|
||||
scheduler.step()
|
||||
|
||||
checkpoint = Checkpoint(
|
||||
optimizer_state_dict=optimizer.state_dict(),
|
||||
scheduler_state_dict=scheduler.state_dict(),
|
||||
state_dict=model.state_dict(),
|
||||
epoch=2,
|
||||
iteration=10,
|
||||
metrics={"loss": [0.3, 0.2, 0.1]}
|
||||
)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
checkpoint.save(tmpdir)
|
||||
loaded = Checkpoint.load(tmpdir)
|
||||
assert loaded.epoch == 2
|
||||
print(f"Rank {rank}: Checkpoint test passed")
|
||||
rank = get_rank()
|
||||
|
||||
if rank == 0:
|
||||
shared_dir = tempfile.mkdtemp()
|
||||
checkpoint.save(shared_dir)
|
||||
else:
|
||||
shared_dir = None
|
||||
|
||||
|
||||
if dist.is_initialized():
|
||||
dir_list = [shared_dir]
|
||||
dist.broadcast_object_list(dir_list, src=0)
|
||||
shared_dir = dir_list[0]
|
||||
|
||||
|
||||
loaded = Checkpoint.load(shared_dir)
|
||||
assert loaded.epoch == 2
|
||||
|
||||
def test_multi_process():
|
||||
spawn_parallel_fn(
|
||||
|
||||
+19
-32
@@ -1,4 +1,3 @@
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
@@ -6,18 +5,6 @@ from khaosz.data.file import save_h5
|
||||
from khaosz.data.dataset import *
|
||||
|
||||
|
||||
def create_h5_dataset(dir_path, data_dict, dataset_name):
|
||||
"""Helper function to create HDF5 dataset for testing"""
|
||||
dataset_path = os.path.join(dir_path, f"{dataset_name}.h5")
|
||||
|
||||
# Convert data_dict to the format expected by save_h5
|
||||
# save_h5 expects a list of tensors for each key
|
||||
tensor_group = {key: [tensor] for key, tensor in data_dict.items()}
|
||||
|
||||
save_h5(dataset_path, tensor_group)
|
||||
|
||||
return dataset_path
|
||||
|
||||
|
||||
def test_dataset_loader_random_paths(base_test_env):
|
||||
"""Test dataset loader with multiple random paths"""
|
||||
@@ -27,23 +14,23 @@ def test_dataset_loader_random_paths(base_test_env):
|
||||
num_files = np.random.randint(2, 5)
|
||||
|
||||
for i in range(num_files):
|
||||
seq_length = np.random.randint(100, 200)
|
||||
seq_length = np.random.randint(200, 400)
|
||||
dummy_data = {
|
||||
"sequence": torch.randint(0, 1000, (seq_length,), dtype=torch.int64),
|
||||
"sequence": [torch.randint(0, 1000, (seq_length,), dtype=torch.int64) for _ in range(10)],
|
||||
}
|
||||
dataset_path = create_h5_dataset(test_dir, dummy_data, f"test_data_{i}")
|
||||
save_h5(test_dir, f"data_{i}", dummy_data)
|
||||
|
||||
# Test loading with multiple paths
|
||||
loaded_dataset = DatasetLoader.load(
|
||||
train_type="seq",
|
||||
load_path=dataset_path,
|
||||
load_path=test_dir,
|
||||
window_size=64,
|
||||
)
|
||||
assert loaded_dataset is not None
|
||||
assert len(loaded_dataset) > 0
|
||||
|
||||
# Test that we can get items without errors
|
||||
for i in range(min(3, len(loaded_dataset))):
|
||||
for i in range(len(loaded_dataset)):
|
||||
item = loaded_dataset[i]
|
||||
assert "input_ids" in item
|
||||
assert "target_ids" in item
|
||||
@@ -59,18 +46,18 @@ def test_dpo_strategy_with_random_data(base_test_env):
|
||||
seq_length = np.random.randint(100, 200)
|
||||
|
||||
dummy_data = {
|
||||
"chosen": torch.randint(0, 1000, (seq_length,), dtype=torch.int64),
|
||||
"rejected": torch.randint(0, 1000, (seq_length,), dtype=torch.int64),
|
||||
"chosen_mask": torch.ones(seq_length, dtype=torch.bool),
|
||||
"rejected_mask": torch.ones(seq_length, dtype=torch.bool)
|
||||
"chosen": [torch.randint(0, 1000, (seq_length,), dtype=torch.int64)],
|
||||
"rejected": [torch.randint(0, 1000, (seq_length,), dtype=torch.int64)],
|
||||
"chosen_mask": [torch.ones(seq_length, dtype=torch.bool)],
|
||||
"rejected_mask": [torch.ones(seq_length, dtype=torch.bool)]
|
||||
}
|
||||
|
||||
dataset_path = create_h5_dataset(test_dir, dummy_data, "dpo_data")
|
||||
save_h5(test_dir, "dpo_data", dummy_data)
|
||||
|
||||
# Load DPO dataset
|
||||
dpo_dataset = DatasetLoader.load(
|
||||
train_type="dpo",
|
||||
load_path=dataset_path,
|
||||
load_path=test_dir,
|
||||
window_size=64,
|
||||
)
|
||||
|
||||
@@ -97,16 +84,16 @@ def test_sft_dataset_with_random_data(base_test_env):
|
||||
seq_length = np.random.randint(100, 200)
|
||||
|
||||
dummy_data = {
|
||||
"sequence": torch.randint(0, 1000, (seq_length,), dtype=torch.int64),
|
||||
"loss_mask": torch.ones(seq_length, dtype=torch.bool)
|
||||
"sequence": [torch.randint(0, 1000, (seq_length,), dtype=torch.int64)],
|
||||
"loss_mask": [torch.ones(seq_length, dtype=torch.bool)]
|
||||
}
|
||||
|
||||
dataset_path = create_h5_dataset(test_dir, dummy_data, "sft_data")
|
||||
save_h5(test_dir, "sft_data", dummy_data)
|
||||
|
||||
# Load SFT dataset
|
||||
sft_dataset = DatasetLoader.load(
|
||||
train_type="sft",
|
||||
load_path=dataset_path,
|
||||
load_path=test_dir,
|
||||
window_size=64,
|
||||
)
|
||||
|
||||
@@ -131,16 +118,16 @@ def test_dataset_with_custom_stride(base_test_env):
|
||||
# Create test data
|
||||
seq_length = 200
|
||||
dummy_data = {
|
||||
"sequence": torch.randint(0, 1000, (seq_length,), dtype=torch.int64),
|
||||
"sequence": [torch.randint(0, 1000, (seq_length,), dtype=torch.int64)],
|
||||
}
|
||||
|
||||
dataset_path = create_h5_dataset(test_dir, dummy_data, "stride_test_data")
|
||||
save_h5(test_dir,"stride_test_data", dummy_data)
|
||||
|
||||
# Test with custom stride
|
||||
custom_stride = 32
|
||||
dataset = DatasetLoader.load(
|
||||
train_type="seq",
|
||||
load_path=dataset_path,
|
||||
load_path=test_dir,
|
||||
window_size=64,
|
||||
stride=custom_stride
|
||||
)
|
||||
@@ -152,7 +139,7 @@ def test_dataset_with_custom_stride(base_test_env):
|
||||
# than with default stride (which equals window size)
|
||||
default_stride_dataset = DatasetLoader.load(
|
||||
train_type="seq",
|
||||
load_path=dataset_path,
|
||||
load_path=test_dir,
|
||||
window_size=64,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user