Team Ai
Modelpublic

baseten/gemma-4-e2b-it-sequence-classification

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes19downloads
modeling_gemma4_sequence.py202 linesDownload Raw Back to root
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