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
This commit is contained in:
2026-08-19 20:55:13 +08:00
parent 398e8a3ea3
commit c79d34eee1
9 changed files with 17 additions and 83 deletions
+3 -8
View File
@@ -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,
+2 -2
View File
@@ -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:
+2 -65
View File
@@ -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",
]
+1 -1
View File
@@ -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:
+5 -3
View File
@@ -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(