Feature Extraction
Transformers
Safetensors
English
remote-sensing
earth-observation
vision
dofa
sentinel-2
multimodal
Instructions to use BiliSakura/DOFA-transformers with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use BiliSakura/DOFA-transformers with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="BiliSakura/DOFA-transformers", device_map="auto")# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("BiliSakura/DOFA-transformers", dtype="auto", device_map="auto") - Notebooks
- Google Colab
- Kaggle
| # 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"] | |