diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loader.py deleted file mode 100644 index e69de29bb..000000000 diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py index 2f96619d3..f255fef5e 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py @@ -19,6 +19,7 @@ from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.loader.utils import ( _normalize_component_type, component_name_to_loader_cls, + get_memory_usage_of_component, ) from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.server_args import ServerArgs @@ -124,14 +125,35 @@ class ComponentLoader(ABC): if isinstance(component, nn.Module): component = component.eval() current_gpu_mem = current_platform.get_available_gpu_memory() + model_size = get_memory_usage_of_component(component) consumed = gpu_mem_before_loading - current_gpu_mem - logger.info( - f"Loaded %s: %s ({source} version). consumed: %.2f GB, avail mem: %.2f GB", - component_name, - component.__class__.__name__, - consumed, - current_gpu_mem, - ) + + # detect component device + try: + component_device = str(next(component.parameters()).device) + is_on_gpu = "cuda" in component_device + except (StopIteration, AttributeError): + is_on_gpu = False + component_device = "unknown" + + if is_on_gpu: + logger.info( + f"Loaded %s: %s ({source} version). model size: %.2f GB, consumed GPU: %.2f GB, avail GPU mem: %.2f GB", + component_name, + component.__class__.__name__, + model_size, + consumed, + current_gpu_mem, + ) + else: + logger.info( + f"Loaded %s: %s ({source} version). model size: %.2f GB, device: %s, avail GPU mem: %.2f GB", + component_name, + component.__class__.__name__, + model_size, + component_device, + current_gpu_mem, + ) return component, consumed def load_native( diff --git a/python/sglang/multimodal_gen/runtime/loader/utils.py b/python/sglang/multimodal_gen/runtime/loader/utils.py index 32aee2bf9..b18603b57 100644 --- a/python/sglang/multimodal_gen/runtime/loader/utils.py +++ b/python/sglang/multimodal_gen/runtime/loader/utils.py @@ -179,5 +179,25 @@ def _list_safetensors_files(model_path: str) -> list[str]: BYTES_PER_GB = 1024**3 + +def get_memory_usage_of_component(module) -> float | None: + """ + returned value is in GB, rounded to 2 decimal digits + """ + if not isinstance(module, nn.Module): + return None + if hasattr(module, "get_memory_footprint"): + usage = module.get_memory_footprint() / BYTES_PER_GB + else: + # manually + param_size = sum(p.numel() * p.element_size() for p in module.parameters()) + buffer_size = sum(b.numel() * b.element_size() for b in module.buffers()) + + total_size_bytes = param_size + buffer_size + usage = total_size_bytes / (1024**3) + + return round(usage, 2) + + # component name -> ComponentLoader class component_name_to_loader_cls: Dict[str, Type[Any]] = {} diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py index 592151637..e8855892f 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py @@ -24,7 +24,10 @@ from sglang.multimodal_gen.configs.sample.teacache import ( TeaCacheParams, WanTeaCacheParams, ) -from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.server_args import ( + ServerArgs, + _sanitize_for_logging, +) from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.perf_logger import RequestMetrics from sglang.multimodal_gen.utils import align_to @@ -291,14 +294,20 @@ class Req: else: target_width = -1 - # Log sampling parameters - debug_str = f"""Sampling params: + # sanitize prompts for info-level logging + sanitized_prompt = _sanitize_for_logging(self.prompt, key_hint="prompt") + sanitized_neg_prompt = _sanitize_for_logging( + self.negative_prompt, key_hint="negative_prompt" + ) + + # log non-sensitive parameters at info level + info_str = f"""Sampling params: width: {target_width} height: {target_height} num_frames: {self.num_frames} fps: {self.fps} - prompt: {self.prompt} - neg_prompt: {self.negative_prompt} + prompt: {sanitized_prompt} + neg_prompt: {sanitized_neg_prompt} seed: {self.seed} infer_steps: {self.num_inference_steps} num_outputs_per_prompt: {self.num_outputs_per_prompt} @@ -310,7 +319,15 @@ class Req: save_output: {self.save_output} output_file_path: {self.output_file_path()} """ # type: ignore[attr-defined] + + # log full prompts at debug level only (for debugging purposes) + debug_str = f"""Full prompts: + prompt: {self.prompt} + neg_prompt: {self.negative_prompt} + """ + if not self.suppress_logs: + logger.info(info_str) logger.debug(debug_str) diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index 72126e0ac..d5fb89a09 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -63,10 +63,15 @@ def _sanitize_for_logging(obj: Any, key_hint: str | None = None) -> Any: - Render torch.Tensor as a compact summary; if key name is 'scaling_factor', include stats. - Dataclasses are expanded to dicts and sanitized recursively. - Callables/functions are rendered as their qualified name. + - Redact sensitive fields like 'prompt' and 'negative_prompt' (only show length). - Fallback to str(...) for unknown types. """ # Handle simple types quickly if obj is None or isinstance(obj, (str, int, float, bool)): + # redact sensitive prompt fields + if key_hint in ("prompt", "negative_prompt"): + if isinstance(obj, str): + return f"" return obj # Enum -> value for readability @@ -134,7 +139,7 @@ def _sanitize_for_logging(obj: Any, key_hint: str | None = None) -> Any: # Sequences/Sets -> list if isinstance(obj, (list, tuple, set)): - return [_sanitize_for_logging(x) for x in obj] + return [_sanitize_for_logging(x, key_hint=key_hint) for x in obj] # Functions / Callables -> qualified name try: