Add transformers 5.x model loading path for Florence2

Uses local Florence2 modules with accelerate-based weight loading and
a custom Florence2Processor for transformers >= 5.0.0, while keeping
the existing loading path for older versions.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Gang Wen
2026-04-07 08:15:29 +02:00
co-authored by Claude Opus 4.6
parent c8b2ec0b59
commit 810e34d859
+68 -7
View File
@@ -6,6 +6,8 @@ import matplotlib.pyplot as plt
import matplotlib.patches as patches
import colorsys
from transformers.dynamic_module_utils import get_imports
import transformers
from packaging import version
import comfy.model_management
from .imagefunc import *
@@ -43,22 +45,81 @@ def fixed_get_imports(filename) -> list[str]:
pass
return imports
def load_model(version):
def _load_model_v5(model_path, attention, dtype):
"""Load Florence2 model for transformers >= 5.0.0"""
from florence2_models.modeling_florence2 import Florence2ForConditionalGeneration, Florence2Config
from transformers import CLIPImageProcessor, BartTokenizerFast
from florence2_models.processing_florence2 import Florence2Processor
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
from comfy.utils import load_torch_file
offload_device = comfy.model_management.unet_offload_device()
config = Florence2Config.from_pretrained(model_path)
config._attn_implementation = attention
with init_empty_weights():
model = Florence2ForConditionalGeneration(config)
checkpoint_path = os.path.join(model_path, "model.safetensors")
if not os.path.exists(checkpoint_path):
checkpoint_path = os.path.join(model_path, "pytorch_model.bin")
if os.path.exists(checkpoint_path):
state_dict = load_torch_file(checkpoint_path)
else:
raise FileNotFoundError(f"No model weights found at {model_path}")
key_mapping = {}
if "language_model.model.shared.weight" in state_dict:
key_mapping["language_model.model.encoder.embed_tokens.weight"] = "language_model.model.shared.weight"
key_mapping["language_model.model.decoder.embed_tokens.weight"] = "language_model.model.shared.weight"
for name, param in model.named_parameters():
actual_key = key_mapping.get(name, name)
if actual_key in state_dict:
set_module_tensor_to_device(model, name, offload_device, value=state_dict[actual_key].to(dtype))
else:
print(f"Parameter {name} not found in state_dict.")
model.language_model.tie_weights()
model = model.eval().to(dtype).to(offload_device)
image_processor = CLIPImageProcessor(
do_resize=True,
size={"height": 768, "width": 768},
resample=3,
do_center_crop=False,
do_rescale=True,
rescale_factor=1/255.0,
do_normalize=True,
image_mean=[0.485, 0.456, 0.406],
image_std=[0.229, 0.224, 0.225],
)
image_processor.image_seq_length = 577
tokenizer = BartTokenizerFast.from_pretrained(model_path)
processor = Florence2Processor(image_processor=image_processor, tokenizer=tokenizer)
return model, processor
def load_model(ver):
florence_path = os.path.join(folder_paths.models_dir, "florence2")
os.makedirs(florence_path, exist_ok=True)
model_path = os.path.join(florence_path, version)
model_path = os.path.join(florence_path, ver)
attention = 'sdpa'
if not os.path.exists(model_path):
log(f"Downloading Florence2 {version} model...")
repo_id = fl2_model_repos[version]
log(f"Downloading Florence2 {ver} model...")
repo_id = fl2_model_repos[ver]
from huggingface_hub import snapshot_download
snapshot_download(repo_id=repo_id, local_dir=model_path, ignore_patterns=["*.md", "*.txt"])
if version.parse(transformers.__version__) >= version.parse('5.0.0'):
model, processor = _load_model_v5(model_path, attention, torch.float32)
return (model.to(device), processor)
try:
with patch("transformers.dynamic_module_utils.get_imports", fixed_get_imports):
# model = AutoModelForCausalLM.from_pretrained(model_path, trust_remote_code=True)
model = AutoModelForCausalLM.from_pretrained(model_path, attn_implementation=attention, device_map=device,
torch_dtype=torch.float32, trust_remote_code=True)
processor = AutoProcessor.from_pretrained(model_path, trust_remote_code=True)
@@ -70,10 +131,10 @@ def load_model(version):
except Exception as e:
sys.path.append(model_path)
# Import the Florence modules
if version == 'large-PromptGen-v1.5':
if ver == 'large-PromptGen-v1.5':
from florence2_large.modeling_florence2 import Florence2ForConditionalGeneration
from florence2_large.configuration_florence2 import Florence2Config
elif version == 'base-PromptGen-v1.5':
elif ver == 'base-PromptGen-v1.5':
from florence2_base_ft.modeling_florence2 import Florence2ForConditionalGeneration
from florence2_base_ft.configuration_florence2 import Florence2Config
else: