style: 修改为显式导入

This commit is contained in:
2026-04-04 16:02:49 +08:00
parent 3346c75584
commit b531232a9b
12 changed files with 38 additions and 34 deletions
+7 -4
View File
@@ -58,10 +58,13 @@ def create_train_config(
TrainConfig instance configured for testing
"""
optimizer_fn = lambda m: torch.optim.AdamW(m.parameters(), lr=0.001)
scheduler_fn = lambda optim: SchedulerFactory.create(
optim, "cosine", warmup_steps=10, lr_decay_steps=10, min_rate=0.05
)
def optimizer_fn(m):
return torch.optim.AdamW(m.parameters(), lr=0.001)
def scheduler_fn(optim):
return SchedulerFactory.create(
optim, "cosine", warmup_steps=10, lr_decay_steps=10, min_rate=0.05
)
return TrainConfig(
strategy=strategy,
+12 -6
View File
@@ -1,15 +1,21 @@
import torch
from astrai.config import *
from astrai.trainer import *
from astrai.config.train_config import TrainConfig
from astrai.trainer.schedule import SchedulerFactory
from astrai.trainer.train_callback import TrainCallback
from astrai.trainer.trainer import Trainer
def test_callback_integration(base_test_env, random_dataset):
"""Test that all callbacks are properly integrated"""
optimizer_fn = lambda model: torch.optim.AdamW(model.parameters())
scheduler_fn = lambda optim: SchedulerFactory.create(
optim, "cosine", warmup_steps=10, lr_decay_steps=10, min_rate=0.05
)
def optimizer_fn(model):
return torch.optim.AdamW(model.parameters())
def scheduler_fn(optim):
return SchedulerFactory.create(
optim, "cosine", warmup_steps=10, lr_decay_steps=10, min_rate=0.05
)
train_config = TrainConfig(
model=base_test_env["model"],
+10 -6
View File
@@ -3,18 +3,22 @@ import os
import numpy as np
import torch
from astrai.config import *
from astrai.config.train_config import TrainConfig
from astrai.data.serialization import Checkpoint
from astrai.trainer import *
from astrai.trainer.schedule import SchedulerFactory
from astrai.trainer.trainer import Trainer
def test_early_stopping_simulation(base_test_env, early_stopping_dataset):
"""Simulate early stopping behavior"""
optimizer_fn = lambda model: torch.optim.AdamW(model.parameters())
scheduler_fn = lambda optim: SchedulerFactory.create(
optim, "cosine", warmup_steps=10, lr_decay_steps=10, min_rate=0.05
)
def optimizer_fn(model):
return torch.optim.AdamW(model.parameters())
def scheduler_fn(optim):
return SchedulerFactory.create(
optim, "cosine", warmup_steps=10, lr_decay_steps=10, min_rate=0.05
)
train_config = TrainConfig(
strategy="seq",
+1 -3
View File
@@ -1,9 +1,7 @@
import numpy as np
import torch
from astrai.config import *
from astrai.data.dataset import *
from astrai.trainer.schedule import *
from astrai.trainer.schedule import SchedulerFactory, CosineScheduler, SGDRScheduler
def test_schedule_factory_random_configs():
-1
View File
@@ -1,4 +1,3 @@
from astrai.data.dataset import *
from astrai.trainer import Trainer
# train_config_factory is injected via fixture