diff --git a/skyrl/backends/skyrl_train_backend.py b/skyrl/backends/skyrl_train_backend.py index ca67f2ed5e..bdfad37d14 100644 --- a/skyrl/backends/skyrl_train_backend.py +++ b/skyrl/backends/skyrl_train_backend.py @@ -429,6 +429,18 @@ def _ensure_inference_engines(self): self._render_server.shutdown() self._render_server = None + def initialize_base_inference(self) -> None: + """Start base-model inference before the first training model is created.""" + if self._inference_engines_initialized: + return + self._cfg = _build_skyrl_train_config(self.base_model, self.config) + if not ray.is_initialized(): + logger.info("Initializing Ray for base-model inference") + initialize_ray(self._cfg) + self._colocate_pg = self._create_colocate_pg() if self._cfg.trainer.placement.colocate_all else None + self._create_new_inference_client() + self._inference_engines_initialized = True + def _lora_signature_from(self, lora_config: types.LoraConfig) -> tuple: # Tinker's public LoraConfig only exposes rank + alpha (plus # seed/train_attn/train_mlp/train_unembed) - pending support https://github.com/NovaSky-AI/SkyRL/issues/1632. @@ -497,6 +509,9 @@ def create_model(self, model_id: str, lora_config: types.LoraConfig, model_role: logger.info("Building models.") self._build_policy(PolicyWorker, model_id=model_id) + if self._inference_engines_initialized: + self._dispatch.set_inference_engine_client(self._inference_engine_client) + self.init_weight_sync_state() if is_lora: self._base_lora_signature = self._lora_signature_from(lora_config) elif model_role == "critic": diff --git a/skyrl/tinker/engine.py b/skyrl/tinker/engine.py index a09d9ac5a9..85413fc010 100644 --- a/skyrl/tinker/engine.py +++ b/skyrl/tinker/engine.py @@ -272,6 +272,10 @@ def __init__( if hasattr(self.backend, "set_inference_state_publisher"): self.backend.set_inference_state_publisher(self._write_inference_state_to_db) + is_colocated = bool(config.backend_config.get("trainer.placement.colocate_all", True)) + if not is_colocated and hasattr(self.backend, "initialize_base_inference"): + self.backend.initialize_base_inference() + # Track last cleanup time for periodic stale session cleanup self._last_cleanup_time: float = time.time() @@ -295,6 +299,11 @@ def _write_inference_state_to_db(self, proxy_url: str | None) -> None: row.updated_at = datetime.now(timezone.utc) session.add(row) session.commit() + if proxy_url is not None: + logger.info( + 'SKYRL_DEPLOYMENT_EVENT {"event":"inference_proxy_published","model":"%s"}', + self.config.base_model, + ) @contextmanager def _checkpoint_status_context(self, model_id: str, checkpoint_id: str, checkpoint_type: types.CheckpointType): diff --git a/tests/tinker/skyrl_train/test_text_only_batch_no_inference.py b/tests/tinker/skyrl_train/test_text_only_batch_no_inference.py index 9b7663aed1..d1e8a13e10 100644 --- a/tests/tinker/skyrl_train/test_text_only_batch_no_inference.py +++ b/tests/tinker/skyrl_train/test_text_only_batch_no_inference.py @@ -128,3 +128,25 @@ def shutdown(self): assert fake_self._renderer is None assert render_server.shutdown_called assert fake_self._render_server is None + + +def test_base_inference_initializes_without_training_dispatch(monkeypatch): + cfg = SimpleNamespace(trainer=SimpleNamespace(placement=SimpleNamespace(colocate_all=False))) + calls = [] + fake_self = SimpleNamespace( + _inference_engines_initialized=False, + base_model="model_test", + config=object(), + _cfg=None, + _colocate_pg=None, + _create_new_inference_client=lambda: calls.append("inference"), + ) + monkeypatch.setattr(skyrl_train_backend, "_build_skyrl_train_config", lambda *_args: cfg) + monkeypatch.setattr(skyrl_train_backend.ray, "is_initialized", lambda: True) + + skyrl_train_backend.SkyRLTrainBackend.initialize_base_inference(fake_self) + + assert fake_self._cfg is cfg + assert not hasattr(fake_self, "_dispatch") + assert fake_self._inference_engines_initialized + assert calls == ["inference"]