isLinXu
Pack FlorenceForge source for embedded HF Spaces deployment
e40db0e
Raw
History Blame Contribute Delete
7.53 kB
"""推理引擎模型与设备加载(从 ``inference.py`` 抽出)。"""
from __future__ import annotations
import json
import logging
import os
from pathlib import Path
from typing import Any, Dict, Optional, Union
import torch
import torch.nn as nn
from ..utils.torch_serialization import safe_torch_load
logger = logging.getLogger(__name__)
def resolve_device(device: str) -> torch.device:
if device == "auto":
if torch.cuda.is_available():
device = "cuda"
else:
mps_backend = vars(torch.backends).get("mps")
mps_available = (
mps_backend is not None
and callable(getattr(mps_backend, "is_available", None))
and mps_backend.is_available()
)
device = "mps" if mps_available else "cpu"
return torch.device(device)
class InferenceModelLoader:
"""封装推理模型的设备解析与安全加载。"""
def __init__(
self,
*,
device: torch.device,
compile_model: bool = False,
allow_unsafe_torch_load: bool = False,
model_revision: Optional[str] = None,
) -> None:
self.device = device
self.compile_model = compile_model
self.allow_unsafe_torch_load = allow_unsafe_torch_load
self.model_revision = model_revision
def build_model_config_kwargs(
self,
model_name: str,
revision: Optional[str] = None,
) -> Dict[str, Any]:
config_kwargs: Dict[str, Any] = {
"model_name": model_name,
"device": str(self.device),
"use_lora": False,
}
effective_revision = self.model_revision or revision
if effective_revision:
config_kwargs["revision"] = effective_revision
return config_kwargs
def load_torch_file(self, model_identifier: str) -> nn.Module:
allow_unsafe = self.allow_unsafe_torch_load or (
os.environ.get("FLORENCE_FORGE_ALLOW_UNSAFE_TORCH_LOAD") == "1"
)
try:
loaded = safe_torch_load(
model_identifier,
map_location=self.device,
context="Inference model",
)
except Exception as safe_exc:
if not allow_unsafe:
raise ValueError(
"安全加载本地 Torch 文件失败。FlorenceForge 默认使用 "
"torch.load(weights_only=True);如果该文件是可信来源的整模型 "
"pickle,请传入 allow_unsafe_torch_load=True 或设置 "
"FLORENCE_FORGE_ALLOW_UNSAFE_TORCH_LOAD=1。"
) from safe_exc
logger.warning(
"正在使用 weights_only=False 加载本地 Torch 文件。"
"这会执行 pickle 反序列化,只应对可信文件启用。"
)
try:
loaded = torch.load(
model_identifier,
map_location=self.device,
weights_only=False,
)
except TypeError:
loaded = torch.load(model_identifier, map_location=self.device)
if not isinstance(loaded, nn.Module):
raise TypeError(
f"本地 Torch 文件加载结果是 {type(loaded).__name__},不是 nn.Module。"
"如果这是 state_dict,请先构建模型结构并传入模型实例。"
)
return loaded
def load(self, model: Union[nn.Module, str, Path]) -> nn.Module:
if isinstance(model, nn.Module):
loaded_model = model
else:
model_identifier = str(model)
model_path = Path(model_identifier)
if model_path.suffix in [".pt", ".pth"] and model_path.is_file():
logger.info("尝试加载本地Torch模型文件: %s", model_identifier)
try:
loaded_model = torch.jit.load(model_identifier, map_location=self.device)
logger.info("TorchScript模型加载成功")
except Exception:
logger.info("TorchScript加载失败,尝试安全加载PyTorch模型文件")
loaded_model = self.load_torch_file(model_identifier)
logger.info("PyTorch模型文件加载成功")
else:
try:
from ..core.model import Florence2MultiTaskModel, ModelConfig
if model_path.is_dir() and (model_path / "adapter_config.json").exists():
logger.info("检测到本地LoRA模型: %s", model_identifier)
with open(model_path / "adapter_config.json", "r", encoding="utf-8") as handle:
adapter_config = json.load(handle)
base_model_name = adapter_config.get(
"base_model_name_or_path", "microsoft/Florence-2-base"
)
config = ModelConfig(
**self.build_model_config_kwargs(
base_model_name,
revision=adapter_config.get("revision"),
)
)
loaded_model = Florence2MultiTaskModel.load_pretrained(
model_identifier,
config=config,
is_peft_model=True,
)
logger.info("LoRA模型加载成功")
else:
logger.info("尝试加载Hugging Face模型: %s", model_identifier)
config = ModelConfig(**self.build_model_config_kwargs(model_identifier))
loaded_model = Florence2MultiTaskModel(config)
loaded_model.load()
logger.info("Hugging Face模型加载成功")
except ImportError as exc:
logger.error("无法导入核心模型组件: %s", exc)
raise ValueError(
"加载Hugging Face模型需要 `florence_forge.core.model` 支持。"
) from exc
except Exception as exc:
logger.error("加载Hugging Face模型 '%s' 失败: %s", model_identifier, exc)
raise ValueError(
"无法加载模型。请检查路径或模型ID是否正确,以及是否需要网络连接。"
) from exc
if not hasattr(loaded_model, "eval"):
raise TypeError(
f"加载结果 {type(loaded_model).__name__} 不支持 eval(),无法用于推理"
)
if hasattr(loaded_model, "to") and "Florence2MultiTaskModel" in str(
loaded_model.__class__
):
loaded_model = loaded_model.to(self.device)
elif hasattr(loaded_model, "to"):
loaded_model = loaded_model.to(self.device)
else:
logger.warning("模型 %s 不支持.to()方法,跳过设备移动", type(loaded_model))
if self.compile_model and hasattr(torch, "compile"):
try:
loaded_model = torch.compile(loaded_model)
logger.info("模型编译完成")
except Exception as exc:
logger.warning("模型编译失败: %s", exc)
return loaded_model