diff --git a/nemo_curator/stages/video/caption/caption_enhancement.py b/nemo_curator/stages/video/caption/caption_enhancement.py index dbca98d8dc..ff0922622a 100644 --- a/nemo_curator/stages/video/caption/caption_enhancement.py +++ b/nemo_curator/stages/video/caption/caption_enhancement.py @@ -67,7 +67,7 @@ def __post_init__(self) -> None: self.prompt_text, ) - def setup(self, worker_metadata: WorkerMetadata | None = None) -> None: # noqa: ARG002 + def _initialize_model(self) -> None: if self.model_variant == "qwen": self.model = QwenLM( model_dir=self.model_dir, @@ -81,8 +81,13 @@ def setup(self, worker_metadata: WorkerMetadata | None = None) -> None: # noqa: self.model.setup() def setup_on_node(self, node_info: NodeInfo, worker_metadata: WorkerMetadata) -> None: # noqa: ARG002 - """Download the weights for the QwenLM model on the node.""" + """Download weights and initialize vLLM once per node to avoid torch.compile race conditions.""" QwenLM.download_weights_on_node(self.model_dir) + self._initialize_model() + + def setup(self, worker_metadata: WorkerMetadata | None = None) -> None: # noqa: ARG002 + if not hasattr(self, "model") or self.model is None: + self._initialize_model() def process(self, task: VideoTask) -> VideoTask: video = task.data diff --git a/nemo_curator/stages/video/caption/caption_generation.py b/nemo_curator/stages/video/caption/caption_generation.py index 43a851a286..03f5666c22 100644 --- a/nemo_curator/stages/video/caption/caption_generation.py +++ b/nemo_curator/stages/video/caption/caption_generation.py @@ -51,7 +51,7 @@ def inputs(self) -> tuple[list[str], list[str]]: def outputs(self) -> tuple[list[str], list[str]]: return ["data"], ["clips"] - def setup(self, worker_metadata: WorkerMetadata | None = None) -> None: # noqa: ARG002 + def _initialize_model(self) -> None: if self.model_variant == "qwen": self.model = QwenVL( model_dir=self.model_dir, @@ -68,8 +68,13 @@ def setup(self, worker_metadata: WorkerMetadata | None = None) -> None: # noqa: self.model.setup() def setup_on_node(self, node_info: NodeInfo, worker_metadata: WorkerMetadata) -> None: # noqa: ARG002 - """Download the weights for the QwenVL model on the node.""" + """Download weights and initialize vLLM once per node to avoid torch.compile race conditions.""" QwenVL.download_weights_on_node(self.model_dir) + self._initialize_model() + + def setup(self, worker_metadata: WorkerMetadata | None = None) -> None: # noqa: ARG002 + if not hasattr(self, "model") or self.model is None: + self._initialize_model() def __post_init__(self) -> None: self.resources = Resources(gpus=1)