# Copyright 2024 The DOFA Authors and The HuggingFace Inc. team. """Self-contained DOFA model, config, and dynamic patch embedding.""" from functools import partial from typing import Optional, Union import torch import torch.nn as nn import torch.nn.functional as F import torch.nn.init as init from timm.models.vision_transformer import Block from transformers.configuration_utils import PretrainedConfig as PreTrainedConfig from transformers.modeling_outputs import BaseModelOutputWithPooling, ImageClassifierOutput from transformers.modeling_utils import PreTrainedModel from transformers.processing_utils import Unpack from transformers.utils import TransformersKwargs, logging logger = logging.get_logger(__name__) class DOFAConfig(PreTrainedConfig): model_type = "dofa" def __init__( self, hidden_size=768, num_hidden_layers=12, num_attention_heads=12, intermediate_size=None, hidden_act="gelu", hidden_dropout_prob=0.0, attention_probs_dropout_prob=0.0, initializer_range=0.02, layer_norm_eps=1e-6, image_size=224, patch_size=16, num_channels=3, qkv_bias=True, wv_planes=128, mlp_ratio=4.0, global_pool=True, default_wavelengths=None, default_image_mean=None, default_image_std=None, head_dropout=0.0, num_labels=0, **kwargs, ): super().__init__(**kwargs) self.hidden_size = hidden_size self.num_hidden_layers = num_hidden_layers self.num_attention_heads = num_attention_heads self.hidden_act = hidden_act self.hidden_dropout_prob = hidden_dropout_prob self.attention_probs_dropout_prob = attention_probs_dropout_prob self.initializer_range = initializer_range self.layer_norm_eps = layer_norm_eps self.image_size = image_size self.patch_size = patch_size self.num_channels = num_channels self.qkv_bias = qkv_bias self.wv_planes = wv_planes self.mlp_ratio = mlp_ratio self.global_pool = global_pool self.default_wavelengths = default_wavelengths self.default_image_mean = default_image_mean self.default_image_std = default_image_std self.head_dropout = head_dropout self.num_labels = num_labels if intermediate_size is None: self.intermediate_size = int(hidden_size * mlp_ratio) else: self.intermediate_size = intermediate_size def get_1d_sincos_pos_embed_from_grid_torch(embed_dim, pos): assert embed_dim % 2 == 0 omega = torch.arange(embed_dim // 2, dtype=torch.float32, device=pos.device) omega /= embed_dim / 2.0 omega = 1.0 / 10000**omega pos = pos.reshape(-1) out = torch.einsum("m,d->md", pos, omega) emb_sin = torch.sin(out) emb_cos = torch.cos(out) return torch.cat([emb_sin, emb_cos], dim=1) class TransformerWeightGenerator(nn.Module): def __init__(self, input_dim, output_dim, embed_dim, num_heads=4, num_layers=1): super().__init__() encoder_layer = nn.TransformerEncoderLayer( d_model=input_dim, nhead=num_heads, activation="gelu", norm_first=False, batch_first=False, dropout=False, ) self.transformer_encoder = nn.TransformerEncoder( encoder_layer, num_layers=num_layers, enable_nested_tensor=False ) self.fc_weight = nn.Linear(input_dim, output_dim) self.fc_bias = nn.Linear(input_dim, embed_dim) self.wt_num = 128 self.weight_tokens = nn.Parameter(torch.empty([self.wt_num, input_dim])) self.bias_token = nn.Parameter(torch.empty([1, input_dim])) torch.nn.init.normal_(self.weight_tokens, std=0.02) torch.nn.init.normal_(self.bias_token, std=0.02) def forward(self, x): pos_wave = x x = torch.cat([self.weight_tokens, pos_wave], dim=0) x = torch.cat([x, self.bias_token], dim=0) transformer_output = self.transformer_encoder(x) weights = self.fc_weight(transformer_output[self.wt_num : -1] + pos_wave) bias = self.fc_bias(transformer_output[-1]) return weights, bias class FCResLayer(nn.Module): def __init__(self, linear_size=128): super().__init__() self.nonlin1 = nn.ReLU(inplace=True) self.nonlin2 = nn.ReLU(inplace=True) self.w1 = nn.Linear(linear_size, linear_size) self.w2 = nn.Linear(linear_size, linear_size) def forward(self, x): y = self.w1(x) y = self.nonlin1(y) y = self.w2(y) y = self.nonlin2(y) return x + y class DOFADynamicPatchEmbed(nn.Module): def __init__(self, wv_planes, inter_dim=128, kernel_size=16, embed_dim=768): super().__init__() self.kernel_size = kernel_size self.wv_planes = wv_planes self.embed_dim = embed_dim self._num_kernel = self.kernel_size * self.kernel_size * self.embed_dim self.inter_dim = inter_dim self.patch_size = (kernel_size, kernel_size) self.weight_generator = TransformerWeightGenerator(wv_planes, self._num_kernel, embed_dim) self.scaler = 0.01 self.fclayer = FCResLayer(wv_planes) self.weight_generator.apply(self._weight_init) self.fclayer.apply(self._weight_init) def _weight_init(self, module): if isinstance(module, nn.Linear): init.xavier_uniform_(module.weight) module.bias.data.fill_(0.01) def forward(self, pixel_values, wavelengths): inplanes = wavelengths.size(0) waves = get_1d_sincos_pos_embed_from_grid_torch(self.wv_planes, wavelengths * 1000) waves = self.fclayer(waves) weight, bias = self.weight_generator(waves) dynamic_weight = weight.view(inplanes, self.kernel_size, self.kernel_size, self.embed_dim) dynamic_weight = dynamic_weight.permute([3, 0, 1, 2]) if bias is not None: bias = bias.view([self.embed_dim]) * self.scaler weights = dynamic_weight * self.scaler dynamic_out = F.conv2d( pixel_values, weights, bias=bias, stride=self.kernel_size, padding=1, dilation=1 ) return dynamic_out.flatten(2).transpose(1, 2), waves def _prepare_wavelengths(wavelengths, pixel_values, default_wavelengths=None): if wavelengths is None: if default_wavelengths is None: raise ValueError( "DOFA requires per-channel wavelengths. Pass `wavelengths` to the model or image processor, " "or set `default_wavelengths` in the model config." ) wavelengths = default_wavelengths if not torch.is_tensor(wavelengths): wavelengths = torch.tensor(wavelengths, dtype=torch.float32, device=pixel_values.device) else: wavelengths = wavelengths.to(device=pixel_values.device, dtype=torch.float32) if wavelengths.dim() == 2 and wavelengths.shape[0] == 1: wavelengths = wavelengths.squeeze(0) elif wavelengths.dim() == 2 and wavelengths.shape[0] == pixel_values.shape[0]: if not torch.allclose(wavelengths[0], wavelengths): raise ValueError("DOFA currently expects identical wavelengths for all items in a batch.") wavelengths = wavelengths[0] if wavelengths.dim() != 1: raise ValueError("`wavelengths` must be a 1D tensor/list with one value per input channel.") if wavelengths.shape[0] != pixel_values.shape[1]: raise ValueError( f"Expected {pixel_values.shape[1]} wavelengths for {pixel_values.shape[1]} channels, " f"but got {wavelengths.shape[0]}." ) return wavelengths class DOFAPreTrainedModel(PreTrainedModel): config_class = DOFAConfig base_model_prefix = "dofa" main_input_name = "pixel_values" input_modalities = ("image",) supports_gradient_checkpointing = True _no_split_modules = ["Block"] _supports_sdpa = False def _init_weights(self, module): super()._init_weights(module) if isinstance(module, DOFAModel): if hasattr(module, "pos_embed"): nn.init.trunc_normal_(module.pos_embed, std=self.config.initializer_range) if hasattr(module, "cls_token"): nn.init.trunc_normal_(module.cls_token, std=self.config.initializer_range) class DOFAModel(DOFAPreTrainedModel): def __init__(self, config: DOFAConfig, add_pooling_layer: bool = True): super().__init__(config) self.config = config image_size = config.image_size if isinstance(config.image_size, int) else config.image_size[0] self.patch_embed = DOFADynamicPatchEmbed( wv_planes=config.wv_planes, inter_dim=config.wv_planes, kernel_size=config.patch_size, embed_dim=config.hidden_size, ) self.num_patches = (image_size // config.patch_size) ** 2 self.cls_token = nn.Parameter(torch.zeros(1, 1, config.hidden_size)) self.pos_embed = nn.Parameter( torch.zeros(1, self.num_patches + 1, config.hidden_size), requires_grad=False ) norm_layer = partial(nn.LayerNorm, eps=config.layer_norm_eps) self.blocks = nn.ModuleList( [ Block( config.hidden_size, config.num_attention_heads, config.mlp_ratio, qkv_bias=config.qkv_bias, norm_layer=norm_layer, ) for _ in range(config.num_hidden_layers) ] ) self.global_pool = config.global_pool if self.global_pool: self.fc_norm = norm_layer(config.hidden_size) self.norm = None else: self.fc_norm = None self.norm = norm_layer(config.hidden_size) self.add_pooling_layer = add_pooling_layer self.post_init() def forward_features(self, pixel_values, wavelengths): patch_tokens, _ = self.patch_embed(pixel_values, wavelengths) patch_tokens = patch_tokens + self.pos_embed[:, 1:, :] cls_token = self.cls_token + self.pos_embed[:, :1, :] cls_tokens = cls_token.expand(pixel_values.shape[0], -1, -1) hidden_states = torch.cat((cls_tokens, patch_tokens), dim=1) for block in self.blocks: hidden_states = block(hidden_states) if self.global_pool: pooled_output = self.fc_norm(hidden_states[:, 1:, :].mean(dim=1)) else: pooled_output = self.norm(hidden_states)[:, 0] return hidden_states, pooled_output def forward( self, pixel_values: Optional[torch.Tensor] = None, wavelengths: Optional[Union[torch.Tensor, list]] = 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 wavelengths = _prepare_wavelengths(wavelengths, pixel_values, self.config.default_wavelengths) last_hidden_state, pooled_output = self.forward_features(pixel_values, wavelengths) if not self.add_pooling_layer: pooled_output = None if not return_dict: return (last_hidden_state, pooled_output) return BaseModelOutputWithPooling(last_hidden_state=last_hidden_state, pooler_output=pooled_output) class DOFAForImageClassification(DOFAPreTrainedModel): def __init__(self, config: DOFAConfig): super().__init__(config) self.num_labels = config.num_labels self.dofa = DOFAModel(config, add_pooling_layer=True) self.head_drop = nn.Dropout(config.head_dropout) 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, wavelengths: Optional[Union[torch.Tensor, list]] = None, labels: Optional[torch.Tensor] = None, return_dict: Optional[bool] = None, **kwargs: Unpack[TransformersKwargs], ) -> ImageClassifierOutput: outputs = self.dofa( pixel_values=pixel_values, wavelengths=wavelengths, return_dict=True, **kwargs, ) pooled_output = self.head_drop(outputs.pooler_output) logits = self.classifier(pooled_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__ = ["DOFAConfig", "DOFAForImageClassification", "DOFAModel", "DOFAPreTrainedModel"]