feat: 增加server, 并且修改测试单元

This commit is contained in:
2026-04-02 15:05:07 +08:00
parent 9f1561afe7
commit 475de51c7d
12 changed files with 616 additions and 99 deletions
+15 -50
View File
@@ -1,63 +1,39 @@
import torch
import numpy as np
from astrai.config import *
from astrai.trainer import *
from astrai.data.dataset import *
from astrai.trainer import Trainer
# train_config_factory is injected via fixture
def test_different_batch_sizes(base_test_env, random_dataset):
def test_different_batch_sizes(base_test_env, random_dataset, train_config_factory):
"""Test training with different batch sizes"""
batch_sizes = [1, 2, 4, 8]
for batch_size in batch_sizes:
schedule_config = CosineScheduleConfig(warmup_steps=10, total_steps=20)
optimizer_fn = lambda model: torch.optim.AdamW(model.parameters())
scheduler_fn = lambda optim: SchedulerFactory.load(optim, schedule_config)
train_config = TrainConfig(
strategy="seq",
train_config = train_config_factory(
model=base_test_env["model"],
dataset=random_dataset,
optimizer_fn=optimizer_fn,
scheduler_fn=scheduler_fn,
ckpt_dir=base_test_env["test_dir"],
n_epoch=1,
test_dir=base_test_env["test_dir"],
device=base_test_env["device"],
batch_size=batch_size,
ckpt_interval=5,
accumulation_steps=1,
max_grad_norm=1.0,
random_seed=np.random.randint(1000),
device_type=base_test_env["device"],
)
assert train_config.batch_size == batch_size
def test_gradient_accumulation(base_test_env, random_dataset):
def test_gradient_accumulation(base_test_env, random_dataset, train_config_factory):
"""Test training with different gradient accumulation steps"""
accumulation_steps_list = [1, 2, 4]
for accumulation_steps in accumulation_steps_list:
schedule_config = CosineScheduleConfig(warmup_steps=10, total_steps=20)
optimizer_fn = lambda model: torch.optim.AdamW(model.parameters())
scheduler_fn = lambda optim: SchedulerFactory.load(optim, schedule_config)
train_config = TrainConfig(
strategy="seq",
train_config = train_config_factory(
model=base_test_env["model"],
optimizer_fn=optimizer_fn,
scheduler_fn=scheduler_fn,
dataset=random_dataset,
ckpt_dir=base_test_env["test_dir"],
n_epoch=1,
test_dir=base_test_env["test_dir"],
device=base_test_env["device"],
batch_size=2,
ckpt_interval=10,
accumulation_steps=accumulation_steps,
max_grad_norm=1.0,
random_seed=42,
device_type=base_test_env["device"],
)
trainer = Trainer(train_config)
@@ -66,7 +42,7 @@ def test_gradient_accumulation(base_test_env, random_dataset):
assert train_config.accumulation_steps == accumulation_steps
def test_memory_efficient_training(base_test_env, random_dataset):
def test_memory_efficient_training(base_test_env, random_dataset, train_config_factory):
"""Test training with memory-efficient configurations"""
# Test with smaller batch sizes and gradient checkpointing
small_batch_configs = [
@@ -76,24 +52,13 @@ def test_memory_efficient_training(base_test_env, random_dataset):
]
for config in small_batch_configs:
schedule_config = CosineScheduleConfig(warmup_steps=10, total_steps=20)
optimizer_fn = lambda model: torch.optim.AdamW(model.parameters())
scheduler_fn = lambda optim: SchedulerFactory.load(optim, schedule_config)
train_config = TrainConfig(
strategy="seq",
train_config = train_config_factory(
model=base_test_env["model"],
dataset=random_dataset,
optimizer_fn=optimizer_fn,
scheduler_fn=scheduler_fn,
ckpt_dir=base_test_env["test_dir"],
n_epoch=1,
test_dir=base_test_env["test_dir"],
device=base_test_env["device"],
batch_size=config["batch_size"],
ckpt_interval=5,
accumulation_steps=config["accumulation_steps"],
max_grad_norm=1.0,
random_seed=42,
device_type=base_test_env["device"],
)
assert train_config.accumulation_steps == config["accumulation_steps"]