ouzhang57's picture
Upload folder using huggingface_hub (part 10)
dd35111 verified
Raw
History Blame Contribute Delete
2.69 kB
from transformers import AutoProcessor
import torch
from ..core import ModelConfig
from ..core.device.npu_compatible_device import get_device_type
from ..models.hpsv3 import HPSv3Model
from .base import Metric
class HPSv3Metric(Metric):
def __init__(self, model: HPSv3Model):
super().__init__()
self.model = model
@classmethod
def from_pretrained(
cls,
model_config: ModelConfig = ModelConfig(model_id="DiffSynth-Studio/ImageMetrics", origin_file_pattern="HPSv3/model.safetensors"),
processor_config: ModelConfig = ModelConfig(model_id="DiffSynth-Studio/ImageMetrics", origin_file_pattern="HPSv3/"),
torch_dtype: torch.dtype = torch.bfloat16,
device: torch.device = get_device_type(),
score_index: int = 0,
use_special_tokens: bool = True,
max_pixels: int = 256 * 28 * 28,
min_pixels: int = 256 * 28 * 28,
processor_kwargs: dict = None,
vram_limit: float = None,
):
processor_kwargs = processor_kwargs or {}
model_pool = cls.download_and_load_models([model_config], torch_dtype=torch_dtype, device=device, vram_limit=vram_limit)
model = model_pool.fetch_model("image_metrics_hpsv3")
processor_config.download_if_necessary()
processor = AutoProcessor.from_pretrained(processor_config.path, padding_side="right", **processor_kwargs)
if use_special_tokens:
special_tokens = ["<|Reward|>"]
processor.tokenizer.add_special_tokens({"additional_special_tokens": special_tokens})
model.special_token_ids = processor.tokenizer.convert_tokens_to_ids(special_tokens)
model.reward_token = "special"
model.config.tokenizer_padding_side = processor.tokenizer.padding_side
model.config.pad_token_id = processor.tokenizer.pad_token_id
if hasattr(model.config, "text_config"):
model.config.text_config.pad_token_id = processor.tokenizer.pad_token_id
model.rm_head.to(torch.float32)
model = HPSv3Model(
model=model,
processor=processor,
use_special_tokens=use_special_tokens,
max_pixels=max_pixels,
min_pixels=min_pixels,
score_index=score_index,
).eval()
return cls(model)
@torch.no_grad()
def score(self, prompt: str | list[str], images):
scores = self.model(prompt, images)
return self.tensor_to_list(scores)
def compute(self, prompt: str | list[str], images):
return self.score(prompt, images)
def forward(self, prompt: str | list[str], images):
return self.score(prompt, images)