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
+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: