import copy import importlib.util import json import logging import os from typing import Any, Dict, List, Tuple import torch from easydict import EasyDict def get_gpu_capability(): if not torch.cuda.is_available(): return None, None try: return torch.cuda.get_device_capability(0) except Exception as e: logging.warning(f"Failed to get GPU capability: {e}") return None, None def is_fp8_supported_gpu(): major, minor = get_gpu_capability() if major is None: return False return (major == 8 and minor == 9) or (major >= 9) def is_ada_architecture_gpu(): major, minor = get_gpu_capability() if major is None: return False return major == 8 and minor == 9 def is_module_installed(module_name): try: spec = importlib.util.find_spec(module_name) return spec is not None except ModuleNotFoundError: return False def get_available_ops(op_mapping): """通用的操作可用性检查函数""" available_ops = [] for op_name, module_name in op_mapping.items(): is_available = is_module_installed(module_name) available_ops.append((op_name, is_available)) return available_ops def get_available_quant_ops(): quant_mapping = {"sgl": "sgl_kernel", "vllm": "vllm", "q8f": "q8_kernels", "torchao": "torchao"} available_ops = get_available_ops(quant_mapping) # Ada架构GPU优先使用q8f if is_ada_architecture_gpu(): q8f_available = next((op for op in available_ops if op[0] == "q8f" and op[1]), None) if q8f_available: available_ops.remove(q8f_available) available_ops.insert(0, q8f_available) return available_ops def get_available_attn_ops(): attn_mapping = {"sage_attn2": "sageattention", "flash_attn3": "flash_attn_interface", "flash_attn2": "flash_attn", "torch_sdpa": "torch"} return get_available_ops(attn_mapping) class LightX2VDefaultConfig: """Central default configuration for LightX2V.""" # 分组常量 DEFAULT_ATTENTION_TYPE = "flash_attn3" DEFAULT_QUANTIZATION_SCHEMES = {"dit": "bf16", "t5": "bf16", "clip": "fp16"} DEFAULT_VIDEO_PARAMS = {"height": 480, "width": 832, "length": 81, "fps": 16, "vae_stride": [4, 8, 8], "patch_size": [1, 2, 2]} DEFAULT_CONFIG = { # Model Configuration "model_cls": "wan2.1", "model_path": "", "task": "t2v", "mode": "infer", # Inference Parameters "infer_steps": 40, "seed": 42, "sample_guide_scale": 5.0, "sample_shift": 5, "enable_cfg": True, "prompt": "", "negative_prompt": "", # Video Parameters "target_height": DEFAULT_VIDEO_PARAMS["height"], "target_width": DEFAULT_VIDEO_PARAMS["width"], "target_video_length": DEFAULT_VIDEO_PARAMS["length"], "fps": DEFAULT_VIDEO_PARAMS["fps"], "vae_stride": DEFAULT_VIDEO_PARAMS["vae_stride"], "patch_size": DEFAULT_VIDEO_PARAMS["patch_size"], # TeaCache "feature_caching": "NoCaching", "teacache_thresh": 0.26, "coefficients": None, "use_ret_steps": False, # Quantization "dit_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["dit"], "t5_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["t5"], "clip_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["clip"], "quant_op": "vllm", "precision_mode": "fp32", "dit_quantized_ckpt": None, "t5_quantized_ckpt": None, "clip_quantized_ckpt": None, "mm_config": {"mm_type": "Default"}, # Memory Optimization "rotary_chunk": False, "rotary_chunk_size": 100, "clean_cuda_cache": False, "torch_compile": False, "attention_type": DEFAULT_ATTENTION_TYPE, "self_attn_1_type": DEFAULT_ATTENTION_TYPE, "cross_attn_1_type": DEFAULT_ATTENTION_TYPE, "cross_attn_2_type": DEFAULT_ATTENTION_TYPE, # CPU Offloading "cpu_offload": False, "offload_granularity": "phase", "offload_ratio": 1.0, "t5_cpu_offload": False, "t5_offload_granularity": "model", "lazy_load": False, "unload_modules": False, # VAE Settings "use_tiny_vae": False, "tiny_vae": False, "tiny_vae_path": None, "use_tiling_vae": False, # Other Settings "lora_path": None, "strength_model": 1.0, "do_mm_calib": False, "parallel_attn_type": None, "parallel_vae": False, "parallel": False, "seq_parallel": False, "cfg_parallel": False, "max_area": False, "use_prompt_enhancer": False, "text_len": 512, "use_31_block": True, } class CoefficientCalculator: """Calculate TeaCache coefficients based on model and resolution.""" COEFFICIENTS = { "t2v": { "1.3b": { "default": [ [ -5.21862437e04, 9.23041404e03, -5.28275948e02, 1.36987616e01, -4.99875664e-02, ], [ 2.39676752e03, -1.31110545e03, 2.01331979e02, -8.29855975e00, 1.37887774e-01, ], ] }, "14b": { "default": [ [ -3.03318725e05, 4.90537029e04, -2.65530556e03, 5.87365115e01, -3.15583525e-01, ], [ -5784.54975374, 5449.50911966, -1811.16591783, 256.27178429, -13.02252404, ], ] }, }, "i2v": { "720p": [ [ 8.10705460e03, 2.13393892e03, -3.72934672e02, 1.66203073e01, -4.17769401e-02, ], [-114.36346466, 65.26524496, -18.82220707, 4.91518089, -0.23412683], ], "480p": [ [ 2.57151496e05, -3.54229917e04, 1.40286849e03, -1.35890334e01, 1.32517977e-01, ], [ -3.02331670e02, 2.23948934e02, -5.25463970e01, 5.87348440e00, -2.01973289e-01, ], ], }, } @classmethod def get_coefficients( cls, task: str, model_size: str, resolution: Tuple[int, int], use_ret_steps: bool, ) -> List[List[float]]: """Get appropriate coefficients for TeaCache.""" if task == "t2v": coeffs = cls.COEFFICIENTS["t2v"].get(model_size, {}).get("default", None) else: # i2v width, height = resolution if height >= 720 or width >= 720: coeffs = cls.COEFFICIENTS["i2v"]["720p"] else: coeffs = cls.COEFFICIENTS["i2v"]["480p"] if coeffs: return coeffs[0] if use_ret_steps else coeffs[1] raise ValueError( f"No coefficients found for task: {task}, model_size: {model_size}, resolution: {resolution}, use_ret_steps: {use_ret_steps}" ) class ModularConfigManager: """Manages modular configuration without presets.""" def __init__(self): self.base_config = copy.deepcopy(LightX2VDefaultConfig.DEFAULT_CONFIG) self._available_attn_ops = None self._available_quant_ops = None def _get_available_ops(self, ops_list: List[Tuple[str, bool]], fallback: str = None) -> List[str]: """从操作列表中提取可用的操作""" available = [op_name for op_name, is_available in ops_list if is_available] if fallback and fallback not in available: available.append(fallback) return available @property def available_attention_types(self) -> List[str]: """Get available attention types.""" if self._available_attn_ops is None: self._available_attn_ops = get_available_attn_ops() return self._get_available_ops(self._available_attn_ops, "torch_sdpa") @property def available_quant_schemes(self) -> List[str]: """Get available quantization schemes.""" if self._available_quant_ops is None: self._available_quant_ops = get_available_quant_ops() return self._get_available_ops(self._available_quant_ops) def _update_from_config(self, updates: Dict, config: Dict, mappings: Dict[str, str]) -> None: """通用配置更新方法""" for config_key, update_key in mappings.items(): if config_key in config: if config_key == "seed" and config[config_key] == -1: continue updates[update_key] = config[config_key] def apply_inference_config(self, config: Dict[str, Any]) -> Dict[str, Any]: """Apply basic inference configuration.""" updates = {} # 基础映射配置 basic_mappings = { "model_cls": "model_cls", "model_path": "model_path", "task": "task", "infer_steps": "infer_steps", "seed": "seed", "sample_shift": "sample_shift", "height": "target_height", "width": "target_width", "video_length": "target_video_length", "fps": "fps", "video_duration": "video_duration", "resize_mode": "resize_mode", "denoising_step_list": "denoising_step_list", "use_31_block": "use_31_block", "prev_frame_length": "prev_frame_length", } self._update_from_config(updates, config, basic_mappings) # CFG特殊处理 if "cfg_scale" in config: updates["sample_guide_scale"] = config["cfg_scale"] updates["enable_cfg"] = config["cfg_scale"] != 1.0 # 注意力类型配置 attention_type = config.get("attention_type", LightX2VDefaultConfig.DEFAULT_ATTENTION_TYPE) for attn_key in ["attention_type", "self_attn_1_type", "cross_attn_1_type", "cross_attn_2_type"]: updates[attn_key] = attention_type # TinyVAE配置 if config.get("use_tiny_vae", False): updates.update({"use_tiny_vae": True, "tiny_vae": True, "tiny_vae_path": os.path.join(config["model_path"], "taew2_1.pth")}) return updates def apply_teacache_config(self, config: Dict[str, Any], model_info: Dict[str, Any]) -> Dict[str, Any]: """Apply TeaCache configuration.""" updates = {} if config.get("enable", False): updates["feature_caching"] = "Tea" updates["teacache_thresh"] = config.get("threshold", 0.26) updates["use_ret_steps"] = config.get("use_ret_steps", False) task = model_info.get("task", "t2v") model_size = "14b" if "14b" in model_info.get("model_cls", "") else "1.3b" resolution = ( model_info.get("target_width", 832), model_info.get("target_height", 480), ) coeffs = CoefficientCalculator.get_coefficients(task, model_size, resolution, updates["use_ret_steps"]) updates["coefficients"] = coeffs else: updates["feature_caching"] = "NoCaching" return updates def _get_mm_type(self, dit_scheme: str, quant_backend: str) -> str: """获取mm_type配置""" if dit_scheme == "bf16": return "Default" base_pattern = f"W-{dit_scheme}-channel-sym-A-{dit_scheme}-channel-sym-dynamic" if quant_backend == "vllm": return f"{base_pattern}-Vllm" elif quant_backend == "sgl": suffix = "-Sgl-ActVllm" if dit_scheme == "int8" else "-Sgl" return f"{base_pattern}{suffix}" elif quant_backend == "q8f": return f"{base_pattern}-Q8F" elif quant_backend == "torchao": return f"{base_pattern}-Torchao" else: return "Default" def apply_quantization_config(self, config: Dict[str, Any], model_path: str) -> Dict[str, Any]: """Apply quantization configuration.""" updates = {} defaults = LightX2VDefaultConfig.DEFAULT_QUANTIZATION_SCHEMES # 获取量化方案 dit_scheme = config.get("dit_quant_scheme", defaults["dit"]) t5_scheme = config.get("t5_quant_scheme", defaults["t5"]) clip_scheme = config.get("clip_quant_scheme", defaults["clip"]) quant_backend = config.get("quant_op", "vllm") updates.update( { "dit_quant_scheme": dit_scheme, "t5_quant_scheme": t5_scheme, "clip_quant_scheme": clip_scheme, "t5_quantized": t5_scheme != defaults["t5"], "clip_quantized": clip_scheme != defaults["clip"], } ) # 设置检查点路径 if dit_scheme != defaults["dit"]: updates["dit_quantized_ckpt"] = os.path.join(model_path, dit_scheme) if t5_scheme != defaults["t5"]: t5_path = os.path.join(model_path, t5_scheme) updates["t5_quantized_ckpt"] = os.path.join(t5_path, f"models_t5_umt5-xxl-enc-{t5_scheme}.pth") if clip_scheme != defaults["clip"]: clip_path = os.path.join(model_path, clip_scheme) updates["clip_quantized_ckpt"] = os.path.join(clip_path, f"clip-{clip_scheme}.pth") # 特殊后端处理 if quant_backend in ["q8f", "torchao"]: backend_suffix = f"int8-{quant_backend}" updates.update({"t5_quant_scheme": backend_suffix, "clip_quant_scheme": backend_suffix}) # 设置mm_config mm_type = self._get_mm_type(dit_scheme, quant_backend) updates["mm_config"] = {"mm_type": mm_type} return updates def apply_memory_optimization(self, config: Dict[str, Any]) -> Dict[str, Any]: """Apply memory optimization settings.""" updates = {} level = config.get("optimization_level", "none") # 级别映射配置 level_configs = { "medium": {"cpu_offload": True}, "high": {"cpu_offload": True, "rotary_chunk": True, "t5_cpu_offload": True, "t5_offload_granularity": "model"}, "extreme": { "cpu_offload": True, "rotary_chunk": True, "clean_cuda_cache": True, "t5_cpu_offload": True, "t5_offload_granularity": "block", "lazy_load": True, "unload_modules": True, }, } # 应用级别配置 if level in level_configs: updates.update(level_configs[level]) # 直接配置项映射 direct_mappings = { "enable_rotary_chunk": "rotary_chunk", "clean_cuda_cache": "clean_cuda_cache", "cpu_offload": "cpu_offload", "lazy_load": "lazy_load", "unload_after_inference": "unload_modules", "use_tiling_vae": "use_tiling_vae", } for config_key, update_key in direct_mappings.items(): if config.get(config_key, False): updates[update_key] = True # 附加配置 if updates.get("rotary_chunk"): updates["rotary_chunk_size"] = config.get("rotary_chunk_size", 100) if updates.get("cpu_offload"): updates.update({"offload_granularity": config.get("offload_granularity", "phase"), "offload_ratio": config.get("offload_ratio", 1.0)}) return updates def _load_model_config(self, model_path: str) -> Dict[str, Any]: """加载模型配置文件""" config_path = os.path.join(model_path, "config.json") if not os.path.exists(config_path): return {} try: with open(config_path, "r") as f: return json.load(f) except Exception as e: logging.warning(f"Failed to load model config: {e}") return {} def build_final_config(self, configs: Dict[str, Dict[str, Any]]) -> EasyDict: """Build final configuration from module configs.""" final_config = copy.deepcopy(self.base_config) # 应用配置模块 config_modules = [("inference", self.apply_inference_config), ("memory", self.apply_memory_optimization)] for module_name, apply_func in config_modules: if module_name in configs: final_config.update(apply_func(configs[module_name])) # 特殊处理的模块 if "teacache" in configs: teacache_updates = self.apply_teacache_config(configs["teacache"], final_config) final_config.update(teacache_updates) if "quantization" in configs: model_path = final_config.get("model_path", "") quant_updates = self.apply_quantization_config(configs["quantization"], model_path) final_config.update(quant_updates) # 加载模型配置 model_config = self._load_model_config(final_config.get("model_path", "")) for key, value in model_config.items(): if key not in final_config or final_config[key] is None: final_config[key] = value return EasyDict(final_config)