baseten/gemma-4-e2b-it-sequence-classification
019
1"""Gemma 4 sequence classifier backed by selected next-token logits.2 3This module is intentionally small: it reuses the Gemma 4 multimodal backbone and4replaces the LM head with a classifier head containing selected token rows.5"""6 7from __future__ import annotations8 9from collections.abc import Sequence10 11import torch12from torch import nn13from transformers.modeling_outputs import SequenceClassifierOutputWithPast14from transformers.models.gemma4.configuration_gemma4 import Gemma4Config15from transformers.models.gemma4.modeling_gemma4 import Gemma4Model, Gemma4PreTrainedModel16 17 18class Gemma4ForSequenceClassification(Gemma4PreTrainedModel):19 """Pool the last text position and score it with selected Gemma 4 token rows."""20 21 config_class = Gemma4Config22 base_model_prefix = "model"23 24 @classmethod25 def _can_set_experts_implementation(cls) -> bool:26 return True27 28 def __init__(29 self,30 config: Gemma4Config,31 source_model: nn.Module | None = None,32 classifier_weight: torch.Tensor | None = None,33 ) -> None:34 super().__init__(config)35 self.num_labels = config.num_labels36 self.model = source_model.model if source_model is not None else Gemma4Model(config)37 self.score = nn.Linear(config.text_config.hidden_size, self.num_labels, bias=False)38 39 if classifier_weight is not None:40 self.score.to(device=classifier_weight.device, dtype=classifier_weight.dtype)41 self.score.weight.data.copy_(classifier_weight)42 43 if source_model is None and classifier_weight is None:44 self.post_init()45 46 @classmethod47 def from_conditional_generation(48 cls,49 model_lm: nn.Module,50 selected_token_ids: Sequence[int],51 labels: Sequence[str],52 ) -> "Gemma4ForSequenceClassification":53 token_ids = torch.tensor(selected_token_ids, device=model_lm.lm_head.weight.device)54 classifier_weight = model_lm.lm_head.weight.index_select(0, token_ids).detach().clone()55 cls.configure_classification_config(model_lm.config, selected_token_ids, labels)56 return cls(model_lm.config, source_model=model_lm, classifier_weight=classifier_weight)57 58 @classmethod59 def configure_classification_config(60 cls,61 config: Gemma4Config,62 selected_token_ids: Sequence[int],63 labels: Sequence[str],64 ) -> None:65 config.num_labels = len(labels)66 config.id2label = {i: label for i, label in enumerate(labels)}67 config.label2id = {label: i for i, label in enumerate(labels)}68 config.classifier_token_ids = {69 label: int(token_id) for label, token_id in zip(labels, selected_token_ids)70 }71 config.architectures = [cls.__name__]72 config.problem_type = "single_label_classification"73 if getattr(config, "pad_token_id", None) is None:74 config.pad_token_id = config.text_config.pad_token_id75 76 def get_input_embeddings(self):77 return self.model.get_input_embeddings()78 79 def set_input_embeddings(self, value):80 self.model.set_input_embeddings(value)81 82 def get_per_layer_input_embeddings(self):83 return self.model.get_per_layer_input_embeddings()84 85 def set_per_layer_input_embeddings(self, value):86 self.model.set_per_layer_input_embeddings(value)87 88 def _last_non_pad_token(89 self,90 logits: torch.Tensor,91 input_ids: torch.LongTensor | None,92 attention_mask: torch.Tensor | None,93 inputs_embeds: torch.FloatTensor | None,94 ) -> torch.Tensor | int:95 batch_size = logits.shape[0]96 if attention_mask is not None:97 token_indices = torch.arange(logits.shape[1], device=logits.device)98 return (attention_mask.to(logits.device) * token_indices).argmax(-1)99 100 pad_token_id = getattr(self.config, "pad_token_id", None)101 if input_ids is not None and pad_token_id is not None:102 token_indices = torch.arange(input_ids.shape[-1], device=logits.device)103 non_pad = input_ids.to(logits.device).ne(pad_token_id)104 return (non_pad * token_indices).argmax(-1)105 106 if batch_size != 1:107 raise ValueError(108 "Cannot infer sequence lengths for a padded batch without a pad token."109 )110 111 if input_ids is None and inputs_embeds is None:112 raise ValueError("Expected input_ids or inputs_embeds.")113 114 return -1115 116 def _apply_final_logit_softcapping(self, logits: torch.Tensor) -> torch.Tensor:117 final_logit_softcapping = self.config.get_text_config().final_logit_softcapping118 if final_logit_softcapping is None:119 return logits120 logits = logits / final_logit_softcapping121 logits = torch.tanh(logits)122 return logits * final_logit_softcapping123 124 def forward(125 self,126 input_ids: torch.LongTensor | None = None,127 pixel_values: torch.FloatTensor | None = None,128 pixel_values_videos: torch.FloatTensor | None = None,129 input_features: torch.FloatTensor | None = None,130 attention_mask: torch.Tensor | None = None,131 input_features_mask: torch.Tensor | None = None,132 position_ids: torch.LongTensor | None = None,133 image_position_ids: torch.LongTensor | None = None,134 video_position_ids: torch.LongTensor | None = None,135 past_key_values=None,136 mm_token_type_ids: torch.LongTensor | None = None,137 inputs_embeds: torch.FloatTensor | None = None,138 labels: torch.LongTensor | None = None,139 use_cache: bool | None = None,140 return_dict: bool | None = None,141 **kwargs,142 ):143 return_dict = return_dict if return_dict is not None else self.config.use_return_dict144 outputs = self.model(145 input_ids=input_ids,146 pixel_values=pixel_values,147 pixel_values_videos=pixel_values_videos,148 input_features=input_features,149 attention_mask=attention_mask,150 input_features_mask=input_features_mask,151 position_ids=position_ids,152 past_key_values=past_key_values,153 mm_token_type_ids=mm_token_type_ids,154 inputs_embeds=inputs_embeds,155 use_cache=use_cache,156 image_position_ids=image_position_ids,157 video_position_ids=video_position_ids,158 return_dict=True,159 **kwargs,160 )161 162 logits = self.score(outputs.last_hidden_state)163 logits = self._apply_final_logit_softcapping(logits)164 sequence_lengths = self._last_non_pad_token(165 logits,166 input_ids,167 attention_mask,168 inputs_embeds,169 )170 pooled_logits = logits[171 torch.arange(logits.shape[0], device=logits.device),172 sequence_lengths,173 ]174 175 loss = None176 if labels is not None:177 labels = labels.to(pooled_logits.device)178 if self.config.problem_type == "regression":179 loss = nn.MSELoss()(pooled_logits.squeeze(), labels.squeeze())180 elif self.config.problem_type == "multi_label_classification":181 loss = nn.BCEWithLogitsLoss()(pooled_logits, labels)182 else:183 loss = nn.CrossEntropyLoss()(184 pooled_logits.view(-1, self.num_labels),185 labels.view(-1),186 )187 188 if not return_dict:189 output = (pooled_logits,) + outputs[1:]190 return ((loss,) + output) if loss is not None else output191 192 return SequenceClassifierOutputWithPast(193 loss=loss,194 logits=pooled_logits,195 past_key_values=outputs.past_key_values,196 hidden_states=outputs.hidden_states,197 attentions=outputs.attentions,198 )199 200 201Gemma4ForSequenceClassification.register_for_auto_class("AutoModelForSequenceClassification")202 