File size: 2,755 Bytes
c8c00f0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 | 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)
|