[diffusion] CI: fix nightly CI (#25241)
This commit is contained in:
@@ -553,6 +553,8 @@ class LTX2SnapshotResidencyStrategy(LTX2TwoStageResidencyStrategy):
|
||||
if self._module_is_on_gpu(target_module):
|
||||
self._record_component_ready("transformer")
|
||||
elif not self._snapshot_strategy.is_ready("transformer"):
|
||||
if self._snapshot_low_vram_mode:
|
||||
self._release_stage2_for_low_vram()
|
||||
self._snapshot_strategy.prefetch_component("transformer", target_module)
|
||||
else:
|
||||
self._record_component_ready("transformer")
|
||||
@@ -623,6 +625,16 @@ class LTX2SnapshotResidencyStrategy(LTX2TwoStageResidencyStrategy):
|
||||
if stage1_param is not None and stage1_param.device.type == "cuda":
|
||||
self._release_module_to_cpu_snapshot("transformer")
|
||||
|
||||
def _release_stage2_for_low_vram(self) -> None:
|
||||
stage2_module = self.pipeline.get_module("transformer_2")
|
||||
stage2_param = (
|
||||
next(stage2_module.parameters(), None)
|
||||
if stage2_module is not None
|
||||
else None
|
||||
)
|
||||
if stage2_param is not None and stage2_param.device.type == "cuda":
|
||||
self._release_module_to_cpu_snapshot("transformer_2")
|
||||
|
||||
def _record_component_ready(self, module_name: str) -> None:
|
||||
self._snapshot_strategy.record_ready(
|
||||
module_name, self.pipeline.get_module(module_name)
|
||||
|
||||
@@ -677,19 +677,37 @@ class ServerArgs(DisaggArgsMixin):
|
||||
)
|
||||
|
||||
if self.strict_ports:
|
||||
requested_ports = []
|
||||
if needs_http:
|
||||
self._require_port(self.port, "HTTP")
|
||||
self._require_port(self.scheduler_port, "Scheduler")
|
||||
requested_ports.append((self.port, "HTTP"))
|
||||
requested_ports.append((self.scheduler_port, "Scheduler"))
|
||||
if self.master_port is not None:
|
||||
self._require_port(self.master_port, "Master")
|
||||
requested_ports.append((self.master_port, "Master"))
|
||||
seen_ports: dict[int, str] = {}
|
||||
for port, name in requested_ports:
|
||||
if port in seen_ports:
|
||||
raise RuntimeError(
|
||||
f"{name} port {port} duplicates {seen_ports[port]} port and "
|
||||
"--strict-ports is enabled."
|
||||
)
|
||||
seen_ports[port] = name
|
||||
self._require_port(port, name)
|
||||
else:
|
||||
settled_ports: set[int] = set()
|
||||
if needs_http:
|
||||
self.port = self.settle_port(self.port)
|
||||
settled_ports.add(self.port)
|
||||
initial_scheduler_port = self.scheduler_port + (
|
||||
random.randint(0, 100) if self.scheduler_port == 5555 else 0
|
||||
)
|
||||
self.scheduler_port = self.settle_port(initial_scheduler_port)
|
||||
self.master_port = self.settle_port(self.master_port, 37)
|
||||
self.scheduler_port = self.settle_port(
|
||||
initial_scheduler_port, avoid=settled_ports
|
||||
)
|
||||
settled_ports.add(self.scheduler_port)
|
||||
if self.master_port is not None:
|
||||
self.master_port = self.settle_port(
|
||||
self.master_port, 37, avoid=settled_ports
|
||||
)
|
||||
|
||||
def _adjust_parallelism(self):
|
||||
sp_unspecified = self.sp_degree is None
|
||||
@@ -1467,16 +1485,21 @@ class ServerArgs(DisaggArgsMixin):
|
||||
return f"tcp://{scheduler_host}:{self.scheduler_port}"
|
||||
|
||||
def settle_port(
|
||||
self, port: int, port_inc: int = 42, max_attempts: int = 100
|
||||
self,
|
||||
port: int,
|
||||
port_inc: int = 42,
|
||||
max_attempts: int = 100,
|
||||
avoid: set[int] | None = None,
|
||||
) -> int:
|
||||
"""
|
||||
Find an available port with retry logic.
|
||||
"""
|
||||
attempts = 0
|
||||
original_port = port
|
||||
avoid = avoid or set()
|
||||
|
||||
while attempts < max_attempts:
|
||||
if is_port_available(port):
|
||||
if port not in avoid and is_port_available(port):
|
||||
if attempts > 0:
|
||||
logger.info(
|
||||
f"Port {original_port} was unavailable, using port {port} instead"
|
||||
|
||||
Reference in New Issue
Block a user