Update generate.py

add gguf support
This commit is contained in:
smthemex
2026-03-04 12:08:33 +08:00
committed by GitHub
parent ecacfbfd02
commit 600306fe78
+71 -20
View File
@@ -87,8 +87,43 @@ class Separator():
return full_audio, vocal_audio, bgm_audio
def load_gguf_checkpoint_vl(gguf_checkpoint_path):
from diffusers.utils import is_gguf_available, is_torch_available
if is_gguf_available() and is_torch_available():
import gguf
from gguf import GGUFReader
from diffusers.quantizers.gguf.utils import SUPPORTED_GGUF_QUANT_TYPES, GGUFParameter
else:
raise ImportError("Please install torch and gguf>=0.10.0 to load a GGUF checkpoint in PyTorch.")
def build_model(Weigths_Path,infer_model_path,version,use_flash_attn):
reader = GGUFReader(gguf_checkpoint_path)
parsed_parameters = {}
for tensor in reader.tensors:
name = tensor.name
quant_type = tensor.tensor_type
# if the tensor is a torch supported dtype do not use GGUFParameter
is_gguf_quant = quant_type not in [gguf.GGMLQuantizationType.F32, gguf.GGMLQuantizationType.F16]
if is_gguf_quant and quant_type not in SUPPORTED_GGUF_QUANT_TYPES:
_supported_quants_str = "\n".join([str(type) for type in SUPPORTED_GGUF_QUANT_TYPES])
raise ValueError(
(
f"{name} has a quantization type: {str(quant_type)} which is unsupported."
"\n\nCurrently the following quantization types are supported: \n\n"
f"{_supported_quants_str}"
"\n\nTo request support for this quantization type please open an issue here: https://github.com/huggingface/diffusers"
)
)
weights = torch.from_numpy(tensor.data.copy())
parsed_parameters[name] = GGUFParameter(weights, quant_type=quant_type) if is_gguf_quant else weights
del reader
gc.collect()
return parsed_parameters
def build_model(Weigths_Path,infer_model_path,version,use_flash_attn,offload_audiolm):
torch.backends.cudnn.enabled = False
curent_dir = os.path.join(current_node_path,"SongGeneration")
RESOLVERS = {
@@ -115,26 +150,43 @@ def build_model(Weigths_Path,infer_model_path,version,use_flash_attn):
cfg.audio_tokenizer_checkpoint_sep=f"Flow1dVAESeparate_{Weigths_Path}/model_septoken/model_2.safetensors"
cfg.conditioners.type_info.QwTextTokenizer.token_path=os.path.join(current_node_path,"SongGeneration/third_party/Qwen2-7B")
cfg.version = version
audiolm = builders.get_lm_model(cfg,version)
checkpoint = torch.load(infer_model_path, map_location='cpu',weights_only=False)
audiolm_state_dict = {k.replace('audiolm.', ''): v for k, v in checkpoint.items() if k.startswith('audiolm')}
cfg.offload_audiolm = offload_audiolm
# #### @tuolaku https://github.com/smthemex/ComfyUI_SongGeneration/issues/37 ##### 暂时取消1.5版本的测试
# # add 1.5 support,test。。。。
# key = "condition_provider.conditioners.type_info.output_proj.weight"
# expected_vocab_size = 151646
# if key in audiolm_state_dict:
# weight = audiolm_state_dict[key]
# if weight.size(0) > expected_vocab_size:
# print(f"[SongGeneration] Trimming {key} from {weight.size(0)} to {expected_vocab_size}")
# audiolm_state_dict[key] = weight[:expected_vocab_size, :]
# #####
audiolm = builders.get_lm_model(cfg,version,offload_audiolm)
if not infer_model_path.endswith(".gguf"):
checkpoint = torch.load(infer_model_path, map_location='cpu',weights_only=False)
audiolm_state_dict = {k.replace('audiolm.', ''): v for k, v in checkpoint.items() if k.startswith('audiolm')}
del checkpoint
audiolm.load_state_dict(audiolm_state_dict, strict=False)
del audiolm_state_dict
else:
from diffusers import GGUFQuantizationConfig
from diffusers.quantizers.gguf import GGUFQuantizer
from diffusers.models.model_loading_utils import load_model_dict_into_meta
g_config = GGUFQuantizationConfig(compute_dtype=torch.float16)
hf_quantizer = GGUFQuantizer(quantization_config=g_config)
hf_quantizer.pre_quantized = True
model_state_dict=load_gguf_checkpoint_vl(infer_model_path)
gc.collect()
hf_quantizer._process_model_before_weight_loading(
audiolm,
device_map=None,
state_dict=model_state_dict
)
load_model_dict_into_meta(
audiolm,
model_state_dict,
hf_quantizer=hf_quantizer,
device_map=None,
dtype=torch.float16,
)
hf_quantizer._process_model_after_weight_loading(audiolm)
del model_state_dict
gc.collect()
audiolm.eval().to(torch.float16)
audiolm.load_state_dict(audiolm_state_dict, strict=False)
audiolm = audiolm.eval().to(torch.float16)
#audiolm = audiolm.cuda().to(torch.float16)
del audiolm_state_dict,checkpoint
return audiolm,cfg
@@ -181,7 +233,6 @@ def infer_stage2(item,audiolm,max_duration,lyric,descriptions,gen_type,cfg,cfg_c
return items
def inference_lowram_step2(model,lyric,descriptions,item,gen_type):
generate_inp = {