import torch import torch.nn as nn import torch.nn.functional as F from transformers import AutoTokenizer, AutoModel class GemmaWrapper(nn.Module): def __init__(self, model_name="google/gemma-3-4b-it", device=None): super().__init__() if device is None: device = "cuda" if torch.cuda.is_available() else "cpu" self.tokenizer = AutoTokenizer.from_pretrained(model_name) self.model = AutoModel.from_pretrained(model_name, attn_implementation='eager').to(device) self.model.eval() for param in self.model.parameters(): param.requires_grad = False def forward(self, texts): inputs = self.tokenizer(texts, return_tensors="pt", padding=True, truncation=True, max_length=512).to(self.model.device) with torch.no_grad(): outputs = self.model(**inputs) last_hidden_state = outputs.last_hidden_state attention_mask = inputs.attention_mask return last_hidden_state, attention_mask.bool()