Update inference.py

This commit is contained in:
smthemex
2025-11-30 10:42:11 +08:00
committed by GitHub
parent 5a46841702
commit f2881db701
+23 -80
View File
@@ -6,6 +6,7 @@ from diffusers import QwenImageEditPipeline
from .hacked_models.scheduler import FlowMatchEulerDiscreteScheduler
from .hacked_models.pipeline import QwenImageEditPipeline
from .hacked_models.models import QwenImageTransformer2DModel
from .hacked_models.pipeline_plus import QwenImageEditPlusPipeline
import sys
from .hacked_models.utils import *
from contextlib import contextmanager
@@ -31,9 +32,11 @@ def temp_patch_module_attr(module_name: str, attr_name: str, new_obj):
pass
def load_model(gguf_path,unet_path,node_path):
plus_mode=False
scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(os.path.join(node_path, "Qwen_Edit_GRAG/Qwen-Image-Edit/scheduler"),torch_dtype=torch.bfloat16,)
if gguf_path is not None:
if "2509" in gguf_path.lower() or "plus" in gguf_path.lower():
plus_mode=True
from diffusers import GGUFQuantizationConfig
with temp_patch_module_attr("diffusers", "QwenImageTransformer2DModel", QwenImageTransformer2DModel):
transformer = QwenImageTransformer2DModel.from_single_file(
@@ -43,20 +46,25 @@ def load_model(gguf_path,unet_path,node_path):
torch_dtype=torch.bfloat16,
)
else:
if "2509" in unet_path.lower() or "plus" in unet_path.lower():
plus_mode=True
print("loading from safetensors")
try:
transformer = QwenImageTransformer2DModel.from_single_file(gguf_path,config=os.path.join(node_path, "Qwen_Edit_GRAG/Qwen-Image-Edit/transformer"),torch_dtype=torch.bfloat16,)
with temp_patch_module_attr("diffusers", "QwenImageTransformer2DModel", QwenImageTransformer2DModel):
transformer = QwenImageTransformer2DModel.from_single_file(gguf_path,config=os.path.join(node_path, "Qwen_Edit_GRAG/Qwen-Image-Edit/transformer"),torch_dtype=torch.bfloat16,)
except:
print("loading from safetensors")
from safetensors.torch import load_file
t_state_dict=load_file(unet_path)
new_dict=replace_key(t_state_dict)
with temp_patch_module_attr("diffusers", "QwenImageTransformer2DModel", QwenImageTransformer2DModel):
unet_config = QwenImageTransformer2DModel.load_config(os.path.join(node_path, "Qwen_Edit_GRAG/Qwen-Image-Edit/transformer/config.json"))
transformer = QwenImageTransformer2DModel.from_config(unet_config).to(torch.bfloat16)
transformer.load_state_dict(new_dict, strict=False)
del t_state_dict,new_dict
pipeline = QwenImageEditPipeline.from_pretrained(os.path.join(node_path, "Qwen_Edit_GRAG/Qwen-Image-Edit"), scheduler = scheduler,vae=None,text_encoder=None,transformer=transformer,torch_dtype=torch.bfloat16,)
from safetensors.torch import load_file
t_state_dict=load_file(unet_path)
new_dict=replace_key(t_state_dict)
with temp_patch_module_attr("diffusers", "QwenImageTransformer2DModel", QwenImageTransformer2DModel):
unet_config = QwenImageTransformer2DModel.load_config(os.path.join(node_path, "Qwen_Edit_GRAG/Qwen-Image-Edit/transformer/config.json"))
transformer = QwenImageTransformer2DModel.from_config(unet_config).to(torch.bfloat16)
transformer.load_state_dict(new_dict, strict=False)
del t_state_dict,new_dict
if plus_mode:
pipeline = QwenImageEditPlusPipeline.from_pretrained(os.path.join(node_path, "Qwen_Edit_GRAG/Qwen-Image-Edit"), scheduler = scheduler,vae=None,text_encoder=None,transformer=transformer,torch_dtype=torch.bfloat16,)
else:
pipeline = QwenImageEditPipeline.from_pretrained(os.path.join(node_path, "Qwen_Edit_GRAG/Qwen-Image-Edit"), scheduler = scheduler,vae=None,text_encoder=None,transformer=transformer,torch_dtype=torch.bfloat16,)
return pipeline
def replace_key(t_state_dict):
@@ -75,9 +83,9 @@ def inference(pipeline,positive,negative,num_inference_steps,seed,true_cfg_scale
"true_cfg_scale": true_cfg_scale,
"negative_prompt": None,
"num_inference_steps": num_inference_steps,
"prompt_embeds": positive[0][0], #pooled_prompt_embeds=positive[0][1].get("pooled_output")
"prompt_embeds": positive[0][0],
"negative_prompt_embeds": negative[0][0],
"image_latents":positive[0][1].get("ref_latents",None) ,
"image_latents":positive[0][1].get("reference_latents",None) ,
"return_dict": False,
"grag_scale":[((512,1.0,1.0),(4096,cond_b,cond_delta))]*60,
}
@@ -91,68 +99,3 @@ def inference(pipeline,positive,negative,num_inference_steps,seed,true_cfg_scale
#parser = argparse.ArgumentParser()
# parser.add_argument("--model_path", type=str, default="Qwen/Qwen-Image-Edit")
# parser.add_argument("--image_path", type=str, required=True)
# parser.add_argument("--edit_prompt", type=str, required=True)
# parser.add_argument("--out_path", type=str, default='./results')
# parser.add_argument("--cond_b", type=float, required=True)
# parser.add_argument("--cond_delta", type=float, required=True)
# args = parser.parse_args()
# scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(
# os.path.join(args.model_path, "scheduler"),
# torch_dtype=torch.bfloat16,
# )
# transformer = QwenImageTransformer2DModel.from_pretrained(
# os.path.join(args.model_path, "transformer"),
# torch_dtype=torch.bfloat16,
# )
# pipeline = QwenImageEditPipeline.from_pretrained(args.model_path, torch_dtype=torch.bfloat16,
# scheduler = scheduler,
# transformer=transformer,
# )
# print("pipeline loaded")
# pipeline.to(torch.bfloat16)
# pipeline.to("cuda")
# pipeline.set_progress_bar_config(disable=None)
# out_path = args.out_path
# os.makedirs(out_path,exist_ok=True)
# print(colored(out_path,color = "green"))
# editing_instruction = args.edit_prompt
# input_image = Image.open(args.image_path).convert('RGB').resize((1024,1024))
# os.makedirs(os.path.join(out_path),exist_ok=True)
# prompt = editing_instruction
# inputs = {
# "image": input_image,
# "prompt": prompt,
# "generator": torch.manual_seed(42),
# "true_cfg_scale": 4.0,
# "negative_prompt": " ",
# "num_inference_steps": 24,
# "return_dict": False,
# "grag_scale":[((512,1.0,1.0),(4096,args.cond_b,args.cond_delta))]*60,
# }
# with torch.inference_mode():
# output = pipeline(**inputs)
# image,x0_images,saved_outputs = output
# image[0].save(os.path.join(out_path,f"{args.image_path.split('/')[-1]}_cond_b-{args.cond_b}_cond_delta-{args.cond_delta}.jpg"))