cck-0702's picture
Clean commit without binary (image) files
c8c00f0
Raw
History Blame Contribute Delete
2.76 kB
import torch
from torch import Tensor, nn
from transformers import (CLIPTextModel, CLIPTokenizer, T5EncoderModel, T5Tokenizer, BitsAndBytesConfig)
class HFEmbedder(nn.Module):
def __init__(self, version: str, max_length: int, is_clip, **hf_kwargs):
super().__init__()
self.is_clip = is_clip
self.max_length = max_length
self.output_key = "pooler_output" if self.is_clip else "last_hidden_state"
# Safely remove 'load_in_8bit' and 'device_map' from hf_kwargs so they don't get passed to __init__
self.is_8bit = hf_kwargs.pop("load_in_8bit", False)
device_map = hf_kwargs.pop("device_map", "cuda")
if self.is_clip:
self.tokenizer: CLIPTokenizer = CLIPTokenizer.from_pretrained(version, max_length=max_length)
self.hf_module: CLIPTextModel = CLIPTextModel.from_pretrained(version, **hf_kwargs)
else:
self.tokenizer: T5Tokenizer = T5Tokenizer.from_pretrained(version, max_length=max_length)
if self.is_8bit:
# Use BitsAndBytesConfig for modern transformers
# Remove torch_dtype conflict if present in kwargs
hf_kwargs.pop("torch_dtype", None)
q_config = BitsAndBytesConfig(load_in_8bit=True)
# Remove torch_dtype conflict if present in hf_kwargs
hf_kwargs.pop("torch_dtype", None)
self.hf_module: T5EncoderModel = T5EncoderModel.from_pretrained(
version,
quantization_config=q_config,
device_map=hf_kwargs.pop("device_map", "cuda"),
**hf_kwargs
)
else:
self.hf_module: T5EncoderModel = T5EncoderModel.from_pretrained(version, **hf_kwargs)
self.hf_module = self.hf_module.eval().requires_grad_(False)
def to(self, *args, **kwargs):
# If loaded in 8-bit, bitsandbytes handles device placement automatically.
# Calling .to() on an 8-bit model will crash, so we skip it.
if self.is_8bit:
return self
return super().to(*args, **kwargs)
def forward(self, text: list[str]) -> Tensor:
batch_encoding = self.tokenizer(
text,
truncation=True,
max_length=self.max_length,
return_length=False,
return_overflowing_tokens=False,
padding="max_length",
return_tensors="pt",
)
outputs = self.hf_module(
input_ids=batch_encoding["input_ids"].to(self.hf_module.device),
attention_mask=None,
output_hidden_states=False,
)
return outputs[self.output_key].to(torch.bfloat16)