feat: refine quantization configuration in bridge.py; enhance model class handling in ModularConfigManager; update lora model retrieval in model_utils.py; add cfg_scale2 parameter in LightX2VInferenceConfig

This commit is contained in:
gaclove
2025-09-05 07:27:27 +00:00
parent 34da467904
commit 57c12ecb91
4 changed files with 44 additions and 16 deletions
+6 -7
View File
@@ -105,7 +105,6 @@ class LightX2VDefaultConfig:
"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"],
"adapter_quant_scheme": DEFAULT_QUANTIZATION_SCHEMES["adapter"],
@@ -301,12 +300,16 @@ class ModularConfigManager:
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
# 注意力类型配置
if config["model_cls"] == "wan2.2_moe":
updates["sample_guide_scale"] = [config["cfg_scale2"], config["cfg_scale2"]]
updates["boundary"] = 0.9
if "wan2.2" in config["model_cls"]:
updates["use_image_encoder"] = False
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
@@ -364,9 +367,6 @@ class ModularConfigManager:
updates = {}
defaults = LightX2VDefaultConfig.DEFAULT_QUANTIZATION_SCHEMES
print("config", config)
# 获取量化方案
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"])
@@ -398,7 +398,6 @@ class ModularConfigManager:
"""Apply memory optimization settings."""
updates = {}
# 直接配置项映射
direct_mappings = {
"enable_rotary_chunk": "rotary_chunk",
"clean_cuda_cache": "clean_cuda_cache",
+23 -6
View File
@@ -45,16 +45,26 @@ def support_model_cls_list() -> List[str]:
]
def scan_loras() -> List[str]:
def get_loras_models(model_path: Path) -> List[str]:
loras = []
base_path = get_model_base_path()
loras_path = base_path / "loras"
if loras_path.exists():
for item in loras_path.iterdir():
if model_path.exists():
for item in model_path.iterdir():
if item.is_file():
if item.suffix.lower() in [".safetensors", ".pt", ".pth", ".ckpt"]:
loras.append(item.name)
return loras
def scan_loras() -> List[str]:
base_path = get_model_base_path()
loras_path = base_path / "loras"
loras = get_loras_models(loras_path)
models_base = folder_paths.models_dir
loras_path = Path(models_base) / "loras"
loras2 = get_loras_models(loras_path)
loras.extend(loras2)
loras.sort()
@@ -82,6 +92,13 @@ def get_lora_full_path(lora_name: str) -> str:
if lora_path.exists():
return str(lora_path)
models_base = folder_paths.models_dir
loras_path = Path(models_base) / "loras" / lora_name
if loras_path.exists():
return str(loras_path)
return ""
+14 -2
View File
@@ -85,6 +85,16 @@ class LightX2VInferenceConfig:
"tooltip": "CFG guidance strength",
},
),
"cfg_scale2": (
"FLOAT",
{
"default": 5.0,
"min": 1.0,
"max": 10.0,
"step": 0.1,
"tooltip": "CFG guidance, lower noise when model cls is Wan2.2 MoE",
},
),
"sample_shift": (
"INT",
{"default": 5, "min": 0, "max": 10, "tooltip": "Sample shift"},
@@ -133,7 +143,7 @@ class LightX2VInferenceConfig:
},
),
"resize_mode": (
["adaptive", "keep_ratio_fixed_area", "fixed_min_area", "fixed_max_area", "fixed_shape"],
["adaptive", "keep_ratio_fixed_area", "fixed_min_area", "fixed_max_area", "fixed_shape", "fixed_min_side"],
{
"default": "adaptive",
"tooltip": "Adaptive resize input image to target aspect ratio",
@@ -143,7 +153,7 @@ class LightX2VInferenceConfig:
"STRING",
{
"default": "720p",
"tooltip": "Fixed shape for input image, e.g., '720p', '480p', when resize_mode is 'keep_ratio_fixed_area'",
"tooltip": "Fixed shape for input image, e.g., '720p', '480p', when resize_mode is 'keep_ratio_fixed_area' or 'fixed_min_side'",
},
),
"segment_length": (
@@ -187,6 +197,7 @@ class LightX2VInferenceConfig:
infer_steps,
seed,
cfg_scale,
cfg_scale2,
sample_shift,
height,
width,
@@ -230,6 +241,7 @@ class LightX2VInferenceConfig:
"infer_steps": infer_steps,
"seed": seed if seed != -1 else np.random.randint(0, 2**32 - 1),
"cfg_scale": cfg_scale,
"cfg_scale2": cfg_scale2,
"sample_shift": sample_shift,
"height": height,
"width": width,