# Copyright 2024 The DeCUR Authors and The HuggingFace Inc. team. """Self-contained DeCUR model and config for trust_remote_code loading.""" from typing import Optional import torch import torch.nn as nn from torchvision import models from transformers.configuration_utils import PretrainedConfig as PreTrainedConfig from transformers.modeling_outputs import BaseModelOutputWithPooling, ImageClassifierOutput from transformers.modeling_utils import PreTrainedModel from transformers.models.segformer.configuration_segformer import SegformerConfig from transformers.models.segformer.modeling_segformer import SegformerModel from transformers.models.vit.configuration_vit import ViTConfig from transformers.models.vit.modeling_vit import ViTModel from transformers.processing_utils import Unpack from transformers.utils import TransformersKwargs, logging from .dat_blocks import DAttentionBaseline logger = logging.get_logger(__name__) class DeCURConfig(PreTrainedConfig): model_type = "decur" def __init__( self, backbone="resnet50", num_channels=13, modality="s2c", hidden_size=2048, image_size=224, patch_size=16, use_rda=False, num_hidden_layers=12, num_attention_heads=6, num_labels=0, **kwargs, ): super().__init__(**kwargs) self.backbone = backbone self.num_channels = num_channels self.modality = modality self.hidden_size = hidden_size self.image_size = image_size self.patch_size = patch_size self.use_rda = use_rda self.num_hidden_layers = num_hidden_layers self.num_attention_heads = num_attention_heads self.num_labels = num_labels def get_vit_config(self) -> ViTConfig: if self.backbone != "vits16": raise ValueError(f"Backbone '{self.backbone}' is not a ViT encoder.") return ViTConfig( hidden_size=self.hidden_size, num_hidden_layers=self.num_hidden_layers, num_attention_heads=self.num_attention_heads, intermediate_size=self.hidden_size * 4, image_size=self.image_size, patch_size=self.patch_size, num_channels=self.num_channels, qkv_bias=True, encoder_stride=self.patch_size, ) def get_segformer_config(self) -> SegformerConfig: if self.backbone not in {"mit_b2", "mit_b5"}: raise ValueError(f"Backbone '{self.backbone}' is not a SegFormer encoder.") depths = {"mit_b2": [3, 4, 6, 3], "mit_b5": [3, 6, 40, 3]}[self.backbone] return SegformerConfig( num_channels=self.num_channels, hidden_sizes=[64, 128, 320, 512], depths=depths, num_attention_heads=[1, 2, 5, 8], sr_ratios=[8, 4, 2, 1], patch_sizes=[7, 3, 3, 3], strides=[4, 2, 2, 2], mlp_ratios=[4, 4, 4, 4], hidden_act="gelu", drop_path_rate=0.1, reshape_last_stage=True, ) class DeCURPreTrainedModel(PreTrainedModel): config: DeCURConfig base_model_prefix = "decur" main_input_name = "pixel_values" input_modalities = ("image",) supports_gradient_checkpointing = False def _build_resnet_encoder(config: DeCURConfig) -> nn.Module: backbone = models.resnet50(weights=None) if config.num_channels != 3: backbone.conv1 = nn.Conv2d(config.num_channels, 64, kernel_size=7, stride=2, padding=3, bias=False) backbone.fc = nn.Identity() return backbone def _build_vit_encoder(config: DeCURConfig) -> ViTModel: return ViTModel(config.get_vit_config(), add_pooling_layer=True) def _build_segformer_encoder(config: DeCURConfig) -> SegformerModel: return SegformerModel(config.get_segformer_config()) def _build_rda_modules() -> tuple[DAttentionBaseline, DAttentionBaseline]: da_l3 = DAttentionBaseline( q_size=(14, 14), kv_size=(14, 14), n_heads=8, n_head_channels=128, n_groups=4, attn_drop=0, proj_drop=0, stride=2, offset_range_factor=-1, use_pe=True, dwc_pe=False, no_off=False, fixed_pe=False, ksize=5, log_cpb=False, ) da_l4 = DAttentionBaseline( q_size=(7, 7), kv_size=(7, 7), n_heads=16, n_head_channels=128, n_groups=8, attn_drop=0, proj_drop=0, stride=1, offset_range_factor=-1, use_pe=True, dwc_pe=False, no_off=False, fixed_pe=False, ksize=3, log_cpb=False, ) return da_l3, da_l4 def _build_encoder(config: DeCURConfig) -> nn.Module: if config.backbone == "resnet50": return _build_resnet_encoder(config) if config.backbone == "vits16": return _build_vit_encoder(config) if config.backbone in {"mit_b2", "mit_b5"}: return _build_segformer_encoder(config) raise ValueError(f"Unsupported backbone '{config.backbone}'") class DeCURModel(DeCURPreTrainedModel): def __init__(self, config: DeCURConfig): super().__init__(config) self.encoder = _build_encoder(config) if config.use_rda and config.backbone == "resnet50": self.da_l3, self.da_l4 = _build_rda_modules() else: self.da_l3 = None self.da_l4 = None self.post_init() def _forward_resnet(self, pixel_values: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: x = self.encoder.conv1(pixel_values) x = self.encoder.bn1(x) x = self.encoder.relu(x) x = self.encoder.maxpool(x) x = self.encoder.layer1(x) x = self.encoder.layer2(x) x = self.encoder.layer3(x) if self.da_l3 is not None: x1, _, _ = self.da_l3(x) x = x + x1 x = self.encoder.layer4(x) if self.da_l4 is not None: x2, _, _ = self.da_l4(x) x = x + x2 last_hidden_state = x.flatten(2).transpose(1, 2) pooler_output = self.encoder.avgpool(x).flatten(1) return last_hidden_state, pooler_output def _forward_vit( self, pixel_values: torch.Tensor, *, interpolate_pos_encoding: bool = True, ) -> tuple[torch.Tensor, torch.Tensor]: outputs = self.encoder( pixel_values=pixel_values, interpolate_pos_encoding=interpolate_pos_encoding, return_dict=True, ) last_hidden_state = outputs.last_hidden_state pooler_output = last_hidden_state[:, 0] return last_hidden_state, pooler_output def _forward_segformer(self, pixel_values: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: outputs = self.encoder(pixel_values=pixel_values, return_dict=True) hidden = outputs.last_hidden_state batch_size, channels, height, width = hidden.shape sequence = hidden.permute(0, 2, 3, 1).reshape(batch_size, height * width, channels) pooler_output = sequence.mean(dim=1) return sequence, pooler_output def forward( self, pixel_values: Optional[torch.Tensor] = None, return_dict: Optional[bool] = None, **kwargs: Unpack[TransformersKwargs], ) -> BaseModelOutputWithPooling: if pixel_values is None: raise ValueError("You must specify `pixel_values`") pixel_values = pixel_values.to(dtype=self.dtype) if return_dict is None: return_dict = self.config.use_return_dict if self.config.backbone == "resnet50": last_hidden_state, pooler_output = self._forward_resnet(pixel_values) elif self.config.backbone == "vits16": interpolate_pos_encoding = kwargs.pop("interpolate_pos_encoding", True) last_hidden_state, pooler_output = self._forward_vit( pixel_values, interpolate_pos_encoding=interpolate_pos_encoding, ) else: last_hidden_state, pooler_output = self._forward_segformer(pixel_values) if not return_dict: return (last_hidden_state, pooler_output) return BaseModelOutputWithPooling(last_hidden_state=last_hidden_state, pooler_output=pooler_output) class DeCURForImageClassification(DeCURPreTrainedModel): def __init__(self, config: DeCURConfig): super().__init__(config) self.decur = DeCURModel(config) self.classifier = ( nn.Linear(config.hidden_size, config.num_labels) if config.num_labels > 0 else nn.Identity() ) self.post_init() def forward( self, pixel_values: Optional[torch.Tensor] = None, labels: Optional[torch.Tensor] = None, return_dict: Optional[bool] = None, **kwargs: Unpack[TransformersKwargs], ) -> ImageClassifierOutput: outputs = self.decur(pixel_values=pixel_values, return_dict=True, **kwargs) logits = self.classifier(outputs.pooler_output) loss = None if labels is not None: loss = self.loss_function(labels, logits, self.config, **kwargs) if not return_dict: output = (logits,) + outputs[1:] return ((loss,) + output) if loss is not None else output return ImageClassifierOutput( loss=loss, logits=logits, hidden_states=outputs.hidden_states, attentions=outputs.attentions, ) __all__ = [ "DeCURConfig", "DeCURForImageClassification", "DeCURModel", "DeCURPreTrainedModel", ]