[diffusion] chore: fix ZMQ binding and model loading for FastWan compatibility (#13978)

Co-authored-by: Han Yu <hyu5@dt-login01.delta.ncsa.illinois.edu>
Co-authored-by: Mick <mickjagger19@icloud.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
This commit is contained in:
yh0903
2025-12-30 18:09:38 +08:00
committed by GitHub
co-authored by Han Yu Mick gemini-code-assist[bot]
parent 7518dc3532
commit 49adb37e37
7 changed files with 31 additions and 28 deletions
@@ -31,7 +31,6 @@ def add_multimodal_gen_serve_args(parser: argparse.ArgumentParser):
def execute_serve_cmd(args: argparse.Namespace, unknown_args: list[str] | None = None): def execute_serve_cmd(args: argparse.Namespace, unknown_args: list[str] | None = None):
"""The entry point for the serve command.""" """The entry point for the serve command."""
server_args = ServerArgs.from_cli_args(args, unknown_args) server_args = ServerArgs.from_cli_args(args, unknown_args)
server_args.post_init_serve()
launch_server(server_args) launch_server(server_args)
if server_args.webui: if server_args.webui:
@@ -145,12 +145,12 @@ class DiffGenerator:
if not sync_scheduler_client.ping(): if not sync_scheduler_client.ping():
raise ConnectionError( raise ConnectionError(
f"Could not connect to remote scheduler at " f"Could not connect to remote scheduler at "
f"{self.server_args.scheduler_endpoint()} with `local mode` as False. " f"{self.server_args.scheduler_endpoint} with `local mode` as False. "
"Please ensure the server is running." "Please ensure the server is running."
) )
logger.info( logger.info(
f"Successfully connected to remote scheduler at " f"Successfully connected to remote scheduler at "
f"{self.server_args.scheduler_endpoint()}." f"{self.server_args.scheduler_endpoint}."
) )
def generate( def generate(
@@ -728,6 +728,7 @@ class TransformerLoader(ComponentLoader):
param_dtype=torch.bfloat16, param_dtype=torch.bfloat16,
reduce_dtype=torch.float32, reduce_dtype=torch.float32,
output_dtype=None, output_dtype=None,
strict=False,
) )
total_params = sum(p.numel() for p in model.parameters()) total_params = sum(p.numel() for p in model.parameters())
@@ -79,6 +79,7 @@ def maybe_load_fsdp_model(
fsdp_inference: bool = False, fsdp_inference: bool = False,
output_dtype: torch.dtype | None = None, output_dtype: torch.dtype | None = None,
pin_cpu_memory: bool = True, pin_cpu_memory: bool = True,
strict: bool = True,
) -> torch.nn.Module: ) -> torch.nn.Module:
""" """
Load the model with FSDP if is training, else load the model without FSDP. Load the model with FSDP if is training, else load the model without FSDP.
@@ -138,7 +139,7 @@ def maybe_load_fsdp_model(
weight_iterator, weight_iterator,
device, device,
default_dtype, default_dtype,
strict=True, strict=strict,
cpu_offload=cpu_offload, cpu_offload=cpu_offload,
param_names_mapping=param_names_mapping_fn, param_names_mapping=param_names_mapping_fn,
) )
@@ -255,9 +256,15 @@ def load_model_from_full_model_state_dict(
for target_param_name, full_tensor in custom_param_sd.items(): for target_param_name, full_tensor in custom_param_sd.items():
meta_sharded_param = meta_sd.get(target_param_name) meta_sharded_param = meta_sd.get(target_param_name)
if meta_sharded_param is None: if meta_sharded_param is None:
raise ValueError( if strict:
f"Parameter {target_param_name} not found in custom model state dict. The hf to custom mapping may be incorrect." raise ValueError(
) f"Parameter {target_param_name} not found in custom model state dict. The hf to custom mapping may be incorrect."
)
else:
logger.warning(
f"Parameter '{target_param_name}' from checkpoint not found in model; skipping. This is expected for optional parameters."
)
continue
if not hasattr(meta_sharded_param, "device_mesh"): if not hasattr(meta_sharded_param, "device_mesh"):
full_tensor = full_tensor.to(device=device, dtype=param_dtype) full_tensor = full_tensor.to(device=device, dtype=param_dtype)
actual_param = param_dict.get(target_param_name) actual_param = param_dict.get(target_param_name)
@@ -49,7 +49,7 @@ class Scheduler:
# Inter-process Communication # Inter-process Communication
self.context = zmq.Context(io_threads=2) self.context = zmq.Context(io_threads=2)
endpoint = server_args.scheduler_endpoint() endpoint = server_args.scheduler_endpoint
if gpu_id == 0: if gpu_id == 0:
# router allocates identify (envelope) for each connection # router allocates identify (envelope) for each connection
self.receiver, actual_endpoint = get_zmq_socket( self.receiver, actual_endpoint = get_zmq_socket(
@@ -70,7 +70,7 @@ class SchedulerClient:
# 100 minute timeout for generation # 100 minute timeout for generation
self.scheduler_socket.setsockopt(zmq.RCVTIMEO, 6000000) self.scheduler_socket.setsockopt(zmq.RCVTIMEO, 6000000)
scheduler_endpoint = self.server_args.scheduler_endpoint() scheduler_endpoint = self.server_args.scheduler_endpoint
self.scheduler_socket.connect(scheduler_endpoint) self.scheduler_socket.connect(scheduler_endpoint)
logger.debug( logger.debug(
f"SchedulerClient connected to backend scheduler at {scheduler_endpoint}" f"SchedulerClient connected to backend scheduler at {scheduler_endpoint}"
@@ -98,7 +98,7 @@ class SchedulerClient:
ping_socket.setsockopt(zmq.LINGER, 0) ping_socket.setsockopt(zmq.LINGER, 0)
ping_socket.setsockopt(zmq.RCVTIMEO, 2000) # 2-second timeout for pings ping_socket.setsockopt(zmq.RCVTIMEO, 2000) # 2-second timeout for pings
endpoint = self.server_args.scheduler_endpoint() endpoint = self.server_args.scheduler_endpoint
try: try:
ping_socket.connect(endpoint) ping_socket.connect(endpoint)
@@ -157,7 +157,7 @@ class AsyncSchedulerClient:
# 100 minute timeout # 100 minute timeout
socket.setsockopt(zmq.RCVTIMEO, 6000000) socket.setsockopt(zmq.RCVTIMEO, 6000000)
endpoint = self.server_args.scheduler_endpoint() endpoint = self.server_args.scheduler_endpoint
socket.connect(endpoint) socket.connect(endpoint)
try: try:
@@ -182,7 +182,7 @@ class AsyncSchedulerClient:
ping_socket.setsockopt(zmq.LINGER, 0) ping_socket.setsockopt(zmq.LINGER, 0)
ping_socket.setsockopt(zmq.RCVTIMEO, 2000) ping_socket.setsockopt(zmq.RCVTIMEO, 2000)
endpoint = self.server_args.scheduler_endpoint() endpoint = self.server_args.scheduler_endpoint
try: try:
ping_socket.connect(endpoint) ping_socket.connect(endpoint)
@@ -220,9 +220,9 @@ class ServerArgs:
# TODO: do not hard code # TODO: do not hard code
master_port: int | None = None master_port: int | None = None
# http server endpoint config, would be ignored in local mode # http server endpoint config
host: str | None = None host: str | None = "127.0.0.1"
port: int | None = None port: int | None = 30000
# TODO: webui and their endpoint, check if webui_port is available. # TODO: webui and their endpoint, check if webui_port is available.
webui: bool = False webui: bool = False
@@ -290,6 +290,8 @@ class ServerArgs:
if self.attention_backend in ["fa3", "fa4"]: if self.attention_backend in ["fa3", "fa4"]:
self.attention_backend = "fa" self.attention_backend = "fa"
# network initialization: port and host
self.port = self.settle_port(self.port)
# Add randomization to avoid race condition when multiple servers start simultaneously # Add randomization to avoid race condition when multiple servers start simultaneously
initial_scheduler_port = self.scheduler_port + random.randint(0, 100) initial_scheduler_port = self.scheduler_port + random.randint(0, 100)
self.scheduler_port = self.settle_port(initial_scheduler_port) self.scheduler_port = self.settle_port(initial_scheduler_port)
@@ -306,6 +308,7 @@ class ServerArgs:
"Failed to load V-MoBA config from %s: %s", self.moba_config_path, e "Failed to load V-MoBA config from %s: %s", self.moba_config_path, e
) )
raise raise
self.check_server_args() self.check_server_args()
# log clean server_args # log clean server_args
@@ -575,12 +578,15 @@ class ServerArgs:
else: else:
return f"http://{self.host}:{self.port}" return f"http://{self.host}:{self.port}"
@property
def scheduler_endpoint(self): def scheduler_endpoint(self):
""" """
Internal endpoint for scheduler Internal endpoint for scheduler.
Prefers the configured host but normalizes localhost -> 127.0.0.1 to avoid ZMQ issues.
""" """
scheduler_host = self.host or "localhost" scheduler_host = self.host
if scheduler_host is None or scheduler_host == "localhost":
scheduler_host = "127.0.0.1"
return f"tcp://{scheduler_host}:{self.scheduler_port}" return f"tcp://{scheduler_host}:{self.scheduler_port}"
def settle_port( def settle_port(
@@ -623,16 +629,6 @@ class ServerArgs:
f"(started from port {original_port})" f"(started from port {original_port})"
) )
def post_init_serve(self):
"""
Post init when in serve mode
"""
if self.host is None:
self.host = "localhost"
if self.port is None:
self.port = 3000
self.port = self.settle_port(self.port)
@classmethod @classmethod
def from_cli_args( def from_cli_args(
cls, args: argparse.Namespace, unknown_args: list[str] | None = None cls, args: argparse.Namespace, unknown_args: list[str] | None = None