From c79d34eee1f4822980a36be19adf949d62f846ab Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Wed, 19 Aug 2026 20:14:53 +0800 Subject: [PATCH] refactor: simplify training and inference interfaces - avoid constructing model_fn more than once when reading config - keep inference package exports focused on public entry points - rename extra strategy arguments to strategy_kwargs --- astrai/__init__.py | 11 ++---- astrai/config/train_config.py | 4 +- astrai/inference/__init__.py | 67 +------------------------------- astrai/trainer/strategy.py | 2 +- astrai/trainer/train_context.py | 8 ++-- scripts/tools/train.py | 2 +- tests/inference/test_cache.py | 2 +- tests/trainer/test_grpo_e2e.py | 2 +- tests/trainer/test_online_e2e.py | 2 +- 9 files changed, 17 insertions(+), 83 deletions(-) diff --git a/astrai/__init__.py b/astrai/__init__.py index de331b8..1d223fc 100644 --- a/astrai/__init__.py +++ b/astrai/__init__.py @@ -17,14 +17,9 @@ from astrai.dataset import ( StoreFactory, ) from astrai.factory import BaseFactory -from astrai.inference import ( - InferenceEngine, - ProtocolHandler, - SamplingPipeline, - get_app, - run_server, - sample, -) +from astrai.inference import InferenceEngine, get_app, run_server, sample +from astrai.inference.network import ProtocolHandler +from astrai.inference.runtime.sample import SamplingPipeline from astrai.logging import setup_logging from astrai.model import ( AutoModel, diff --git a/astrai/config/train_config.py b/astrai/config/train_config.py index 0448333..31330d3 100644 --- a/astrai/config/train_config.py +++ b/astrai/config/train_config.py @@ -70,7 +70,7 @@ class TrainConfig(BaseConfig): rollout_max_tokens (int): Maximum generated tokens per response in rollout. Defaults to 1024. reward_model_fn (Optional[Callable]): Factory for reward model, required for online RL strategies. Defaults to None. executor_kwargs (Dict[str, Any]): Extra kwargs passed to ExecutorFactory.create(). Defaults to {}. - extra_kwargs (Dict[str, Any]): Other arguments. Defaults to {}. + strategy_kwargs (Dict[str, Any]): Extra strategy arguments. Defaults to {}. """ model_fn: Callable[[], nn.Module] @@ -125,7 +125,7 @@ class TrainConfig(BaseConfig): reward_model_fn: Optional[Callable] = None executor_kwargs: Dict[str, Any] = field(default_factory=dict) - extra_kwargs: Dict[str, Any] = field(default_factory=dict) + strategy_kwargs: Dict[str, Any] = field(default_factory=dict) @field_validator("strategy") def _validate_strategy(cls, v: str) -> str: diff --git a/astrai/inference/__init__.py b/astrai/inference/__init__.py index d44a411..3fadee1 100644 --- a/astrai/inference/__init__.py +++ b/astrai/inference/__init__.py @@ -12,45 +12,10 @@ Modules: - engine.py: Facade (InferenceEngine) """ -from astrai.inference.cache import ( - Allocator, - KVCache, - KVStorage, - PagePool, - RadixCache, - ReqToTokenPool, - TaskCacheManager, - page_hash, -) from astrai.inference.engine import InferenceEngine -from astrai.inference.network import ( - AnthropicMessage, - BaseToolParser, - ChatCompletionRequest, - ChatMessage, - FunctionDef, - GenContext, - MessagesRequest, - ProtocolHandler, - SimpleJsonToolParser, - StopChecker, - ToolDef, - ToolParserFactory, - get_app, - run_server, -) -from astrai.inference.network.anthropic import AnthropicResponseBuilder -from astrai.inference.network.openai import OpenAIResponseBuilder +from astrai.inference.network import get_app, run_server from astrai.inference.runtime.executor import Executor -from astrai.inference.runtime.sample import ( - BaseSamplingStrategy, - FrequencyPenaltyStrategy, - SamplingPipeline, - TemperatureStrategy, - TopKStrategy, - TopPStrategy, - sample, -) +from astrai.inference.runtime.sample import sample from astrai.inference.scheduler import InferenceScheduler from astrai.inference.task import STOP, Task, TaskManager, TaskStatus @@ -62,35 +27,7 @@ __all__ = [ "Task", "TaskManager", "TaskStatus", - "Allocator", - "KVCache", - "KVStorage", - "PagePool", - "RadixCache", - "ReqToTokenPool", - "TaskCacheManager", - "page_hash", "sample", - "BaseSamplingStrategy", - "TemperatureStrategy", - "TopKStrategy", - "TopPStrategy", - "FrequencyPenaltyStrategy", - "SamplingPipeline", - "ProtocolHandler", - "StopChecker", - "GenContext", - "BaseToolParser", - "SimpleJsonToolParser", - "ToolParserFactory", - "OpenAIResponseBuilder", - "AnthropicResponseBuilder", - "ChatMessage", - "ChatCompletionRequest", - "FunctionDef", - "ToolDef", - "AnthropicMessage", - "MessagesRequest", "get_app", "run_server", ] diff --git a/astrai/trainer/strategy.py b/astrai/trainer/strategy.py index 480284c..08fef52 100644 --- a/astrai/trainer/strategy.py +++ b/astrai/trainer/strategy.py @@ -184,7 +184,7 @@ class BaseStrategy(ABC): self.executor = kwargs.pop("executor", None) self.moe_aux_loss_coef = kwargs.pop("moe_aux_loss_coef", 0.01) self._moe_metrics: Dict[str, float] = {} - self.extra_kwargs = kwargs + self.strategy_kwargs = kwargs self._rollout_runner = None def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor: diff --git a/astrai/trainer/train_context.py b/astrai/trainer/train_context.py index 64bf2c4..c29f1e6 100644 --- a/astrai/trainer/train_context.py +++ b/astrai/trainer/train_context.py @@ -140,8 +140,10 @@ class TrainContextBuilder: checkpoint.consumed_samples // per_step * per_step ) state.checkpoint = checkpoint - if not state.model_config and hasattr(cfg.model_fn(), "config"): - state.model_config = cfg.model_fn().config.to_dict() + if not state.model_config: + model = cfg.model_fn() + if hasattr(model, "config"): + state.model_config = model.config.to_dict() return state def _create_context( @@ -260,7 +262,7 @@ class TrainContextBuilder: def _create_strategy(self, context: TrainContext, executor: BaseExecutor) -> dict: cfg = self.config - kwargs = dict(cfg.extra_kwargs) + kwargs = dict(cfg.strategy_kwargs) kwargs.setdefault("moe_aux_loss_coef", cfg.moe_aux_loss_coef) if cfg.strategy in ("dpo", "grpo", "online_grpo", "online_dpo"): kwargs["ref_model"] = create_ref_model( diff --git a/scripts/tools/train.py b/scripts/tools/train.py index 9eda7c6..1f9690e 100644 --- a/scripts/tools/train.py +++ b/scripts/tools/train.py @@ -836,7 +836,7 @@ def train( gradient_checkpointing_modules=grad_ckpt_modules, compile_mode=compile_mode, executor_kwargs=executor_kwargs, - extra_kwargs=strategy_kwargs, + strategy_kwargs=strategy_kwargs, neftune_alpha=neftune_alpha, collate_fn=collate_fn, rollout_interval=rollout_interval, diff --git a/tests/inference/test_cache.py b/tests/inference/test_cache.py index 6826ef6..2145c77 100644 --- a/tests/inference/test_cache.py +++ b/tests/inference/test_cache.py @@ -2,7 +2,7 @@ import torch -from astrai.inference import ( +from astrai.inference.cache import ( Allocator, KVStorage, PagePool, diff --git a/tests/trainer/test_grpo_e2e.py b/tests/trainer/test_grpo_e2e.py index c82d515..d44b006 100644 --- a/tests/trainer/test_grpo_e2e.py +++ b/tests/trainer/test_grpo_e2e.py @@ -107,7 +107,7 @@ def test_online_grpo_end_to_end(base_test_env): device_type=device, nprocs=1, parallel_mode="none", - extra_kwargs={"clip_eps": 0.2, "kl_coef": 0.01, "group_size": 2}, + strategy_kwargs={"clip_eps": 0.2, "kl_coef": 0.01, "group_size": 2}, rollout_interval=1, rollout_temperature=1.0, rollout_top_k=0, diff --git a/tests/trainer/test_online_e2e.py b/tests/trainer/test_online_e2e.py index d265e5c..9439ac1 100644 --- a/tests/trainer/test_online_e2e.py +++ b/tests/trainer/test_online_e2e.py @@ -109,7 +109,7 @@ def test_online_dpo_end_to_end(base_test_env): device_type=device, nprocs=1, parallel_mode="none", - extra_kwargs={"beta": 0.1, "group_size": 2}, + strategy_kwargs={"beta": 0.1, "group_size": 2}, rollout_interval=1, rollout_temperature=1.0, rollout_top_k=0,