# Multilabel classification wrapper with focal loss class weighting class MistralForMultiLabel(nn.Module): is_loaded_in_4bit = True def __init__(self, backbone: nn.Module, num_labels: int, hidden_size: int, embed_dim: int): super().__init__() self.backbone = backbone _device = torch.device("cuda" if torch.cuda.is_available() else "cpu") self.projection = nn.Sequential( nn.Linear(embed_dim, hidden_size // 2), nn.GELU(), nn.Linear(hidden_size // 2, hidden_size), ).to(_device) self.dropout = nn.Dropout(0.1).to(_device) self.classifier = nn.Linear(hidden_size, num_labels).to(_device) self.focal_loss = FocalLossWithAlpha(FOCAL_ALPHA_PER_LABEL).to(_device) def gradient_checkpointing_enable(self, gradient_checkpointing_kwargs=None): self.backbone.gradient_checkpointing_enable(gradient_checkpointing_kwargs) def gradient_checkpointing_disable(self): self.backbone.gradient_checkpointing_disable() def forward( self, input_embeds: torch.Tensor, labels: torch.Tensor | None = None, **kwargs, ): B = input_embeds.size(0) projected = self.projection(input_embeds).unsqueeze(1) attn_mask = torch.ones(B, 1, device=input_embeds.device) outputs = self.backbone.base_model.model.model( inputs_embeds=projected, attention_mask=attn_mask, output_hidden_states=True, ) pooled = outputs.hidden_states[-1][:, 0, :] logits = self.classifier(self.dropout(pooled)) loss = self.focal_loss(logits, labels.float()) if labels is not None else None return {"loss": loss, "logits": logits}