Feature Extraction
Transformers
Safetensors
English
remote-sensing
earth-observation
vision
croma
sentinel-1
sentinel-2
multimodal
Instructions to use BiliSakura/CROMA-transformers with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use BiliSakura/CROMA-transformers with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="BiliSakura/CROMA-transformers")# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("BiliSakura/CROMA-transformers", device_map="auto") - Notebooks
- Google Colab
- Kaggle
| # Copyright 2023 The CROMA Authors and The HuggingFace Inc. team. | |
| """Self-contained CROMA model and configuration.""" | |
| from __future__ import annotations | |
| import itertools | |
| import math | |
| from dataclasses import dataclass | |
| from typing import Optional | |
| import torch | |
| from einops import rearrange | |
| from torch import einsum, nn | |
| from transformers.configuration_utils import PretrainedConfig as PreTrainedConfig | |
| from transformers.modeling_outputs import BaseModelOutputWithPooling | |
| from transformers.modeling_utils import PreTrainedModel | |
| from transformers.processing_utils import Unpack | |
| from transformers.utils import TransformersKwargs, logging | |
| logger = logging.get_logger(__name__) | |
| class CromaConfig(PreTrainedConfig): | |
| model_type = "croma" | |
| def __init__( | |
| self, | |
| hidden_size=768, | |
| num_hidden_layers=12, | |
| num_attention_heads=16, | |
| patch_size=8, | |
| image_size=120, | |
| sar_channels=2, | |
| optical_channels=12, | |
| modality="both", | |
| hidden_dropout_prob=0.0, | |
| layer_norm_eps=1e-5, | |
| initializer_range=0.02, | |
| num_patches=None, | |
| **kwargs, | |
| ): | |
| super().__init__(**kwargs) | |
| if image_size % patch_size != 0: | |
| raise ValueError(f"`image_size` ({image_size}) must be divisible by `patch_size` ({patch_size}).") | |
| if modality not in {"both", "sar", "optical"}: | |
| raise ValueError(f"`modality` must be one of 'both', 'sar', or 'optical', got {modality!r}.") | |
| self.hidden_size = hidden_size | |
| self.num_hidden_layers = num_hidden_layers | |
| self.num_attention_heads = num_attention_heads | |
| self.patch_size = patch_size | |
| self.image_size = image_size | |
| self.sar_channels = sar_channels | |
| self.optical_channels = optical_channels | |
| self.modality = modality | |
| self.hidden_dropout_prob = hidden_dropout_prob | |
| self.layer_norm_eps = layer_norm_eps | |
| self.initializer_range = initializer_range | |
| self.num_patches = num_patches if num_patches is not None else (image_size // patch_size) ** 2 | |
| def get_2dalibi(num_heads: int, num_patches: int) -> torch.Tensor: | |
| grid_size = int(math.sqrt(num_patches)) | |
| points = list(itertools.product(range(grid_size), range(grid_size))) | |
| def get_slopes(n): | |
| def get_slopes_power_of_2(n): | |
| start = 2 ** (-2 ** -(math.log2(n) - 3)) | |
| ratio = start | |
| return [start * ratio**i for i in range(n)] | |
| if math.log2(n).is_integer(): | |
| return get_slopes_power_of_2(n) | |
| closest_power_of_2 = 2 ** math.floor(math.log2(n)) | |
| return get_slopes_power_of_2(closest_power_of_2) + get_slopes(2 * closest_power_of_2)[0::2][ | |
| : n - closest_power_of_2 | |
| ] | |
| slopes = torch.tensor(get_slopes(num_heads), dtype=torch.float32).unsqueeze(1) | |
| idxs = [] | |
| for p1 in points: | |
| for p2 in points: | |
| dist = math.sqrt((p1[0] - p2[0]) ** 2 + (p1[1] - p2[1]) ** 2) | |
| idxs.append(dist * slopes * -1) | |
| all_bias = torch.cat(idxs, dim=1) | |
| return all_bias.view(1, num_heads, num_patches, num_patches) | |
| class CromaModelOutput(BaseModelOutputWithPooling): | |
| sar_hidden_states: Optional[torch.FloatTensor] = None | |
| sar_pooler_output: Optional[torch.FloatTensor] = None | |
| optical_hidden_states: Optional[torch.FloatTensor] = None | |
| optical_pooler_output: Optional[torch.FloatTensor] = None | |
| joint_hidden_states: Optional[torch.FloatTensor] = None | |
| joint_pooler_output: Optional[torch.FloatTensor] = None | |
| class CromaFeedForward(nn.Module): | |
| def __init__(self, config: CromaConfig, mult: int = 4): | |
| super().__init__() | |
| inner_dim = int(config.hidden_size * mult) | |
| self.net = nn.Sequential( | |
| nn.Linear(config.hidden_size, inner_dim), | |
| nn.GELU(), | |
| nn.Dropout(config.hidden_dropout_prob), | |
| nn.Linear(inner_dim, config.hidden_size), | |
| ) | |
| self.input_norm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) | |
| def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: | |
| hidden_states = self.input_norm(hidden_states) | |
| return self.net(hidden_states) | |
| class CromaAttention(nn.Module): | |
| def __init__(self, config: CromaConfig): | |
| super().__init__() | |
| self.num_attention_heads = config.num_attention_heads | |
| self.attention_head_size = config.hidden_size // config.num_attention_heads | |
| self.scale = self.attention_head_size**-0.5 | |
| self.to_qkv = nn.Linear(config.hidden_size, config.hidden_size * 3, bias=False) | |
| self.to_out = nn.Linear(config.hidden_size, config.hidden_size) | |
| self.input_norm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) | |
| self.dropout = nn.Dropout(config.hidden_dropout_prob) | |
| def forward(self, hidden_states: torch.Tensor, relative_position_bias: torch.Tensor) -> torch.Tensor: | |
| hidden_states = self.input_norm(hidden_states) | |
| query, key, value = self.to_qkv(hidden_states).chunk(3, dim=-1) | |
| query, key, value = map( | |
| lambda tensor: rearrange(tensor, "b n (h d) -> b h n d", h=self.num_attention_heads), | |
| (query, key, value), | |
| ) | |
| attention_scores = einsum("b h i d, b h j d -> b h i j", query, key) * self.scale | |
| attention_scores = attention_scores + relative_position_bias | |
| attention_probs = attention_scores.softmax(dim=-1) | |
| attention_probs = self.dropout(attention_probs) | |
| context = einsum("b h i j, b h j d -> b h i d", attention_probs, value) | |
| context = rearrange(context, "b h n d -> b n (h d)") | |
| return self.to_out(context) | |
| class CromaCrossAttention(nn.Module): | |
| def __init__(self, config: CromaConfig): | |
| super().__init__() | |
| self.num_attention_heads = config.num_attention_heads | |
| self.attention_head_size = config.hidden_size // config.num_attention_heads | |
| self.scale = self.attention_head_size**-0.5 | |
| self.to_q = nn.Linear(config.hidden_size, config.hidden_size, bias=False) | |
| self.to_k = nn.Linear(config.hidden_size, config.hidden_size, bias=False) | |
| self.to_v = nn.Linear(config.hidden_size, config.hidden_size, bias=False) | |
| self.to_out = nn.Linear(config.hidden_size, config.hidden_size) | |
| self.input_norm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) | |
| self.dropout = nn.Dropout(config.hidden_dropout_prob) | |
| def forward( | |
| self, | |
| hidden_states: torch.Tensor, | |
| context: torch.Tensor, | |
| relative_position_bias: torch.Tensor, | |
| ) -> torch.Tensor: | |
| hidden_states = self.input_norm(hidden_states) | |
| context = self.input_norm(context) | |
| query = self.to_q(hidden_states) | |
| key = self.to_k(context) | |
| value = self.to_v(context) | |
| query, key, value = map( | |
| lambda tensor: rearrange(tensor, "b n (h d) -> b h n d", h=self.num_attention_heads), | |
| (query, key, value), | |
| ) | |
| attention_scores = einsum("b h i d, b h j d -> b h i j", query, key) * self.scale | |
| attention_scores = attention_scores + relative_position_bias | |
| attention_probs = attention_scores.softmax(dim=-1) | |
| attention_probs = self.dropout(attention_probs) | |
| context = einsum("b h i j, b h j d -> b h i d", attention_probs, value) | |
| context = rearrange(context, "b h n d -> b n (h d)") | |
| return self.to_out(context) | |
| class CromaEncoder(nn.Module): | |
| def __init__(self, config: CromaConfig, depth: int, final_norm: bool = True): | |
| super().__init__() | |
| self.layers = nn.ModuleList( | |
| [ | |
| nn.ModuleList([CromaAttention(config), CromaFeedForward(config)]) | |
| for _ in range(depth) | |
| ] | |
| ) | |
| self.norm_out = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) if final_norm else None | |
| def forward(self, hidden_states: torch.Tensor, relative_position_bias: torch.Tensor) -> torch.Tensor: | |
| for self_attn, ffn in self.layers: | |
| hidden_states = self_attn(hidden_states, relative_position_bias) + hidden_states | |
| hidden_states = ffn(hidden_states) + hidden_states | |
| if self.norm_out is not None: | |
| return self.norm_out(hidden_states) | |
| return hidden_states | |
| class CromaCrossEncoder(nn.Module): | |
| def __init__(self, config: CromaConfig, depth: int): | |
| super().__init__() | |
| self.layers = nn.ModuleList( | |
| [ | |
| nn.ModuleList([CromaAttention(config), CromaCrossAttention(config), CromaFeedForward(config)]) | |
| for _ in range(depth) | |
| ] | |
| ) | |
| self.norm_out = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) | |
| def forward( | |
| self, | |
| hidden_states: torch.Tensor, | |
| context: torch.Tensor, | |
| relative_position_bias: torch.Tensor, | |
| ) -> torch.Tensor: | |
| for self_attn, cross_attn, ffn in self.layers: | |
| hidden_states = self_attn(hidden_states, relative_position_bias) + hidden_states | |
| hidden_states = cross_attn(hidden_states, context, relative_position_bias) + hidden_states | |
| hidden_states = ffn(hidden_states) + hidden_states | |
| return self.norm_out(hidden_states) | |
| class CromaViTEncoder(nn.Module): | |
| def __init__(self, config: CromaConfig, depth: int, in_channels: int): | |
| super().__init__() | |
| self.patch_size = config.patch_size | |
| pixels_per_patch = config.patch_size * config.patch_size * in_channels | |
| self.linear_input = nn.Linear(pixels_per_patch, config.hidden_size) | |
| self.transformer = CromaEncoder(config, depth=depth) | |
| def forward(self, pixel_values: torch.Tensor, relative_position_bias: torch.Tensor) -> torch.Tensor: | |
| hidden_states = rearrange( | |
| pixel_values, | |
| "b c (h i) (w j) -> b (h w) (c i j)", | |
| i=self.patch_size, | |
| j=self.patch_size, | |
| ) | |
| hidden_states = self.linear_input(hidden_states) | |
| return self.transformer(hidden_states, relative_position_bias) | |
| class CromaGapHead(nn.Module): | |
| def __init__(self, config: CromaConfig): | |
| super().__init__() | |
| self.net = nn.Sequential( | |
| nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps), | |
| nn.Linear(config.hidden_size, 4 * config.hidden_size), | |
| nn.GELU(), | |
| nn.Linear(4 * config.hidden_size, config.hidden_size), | |
| ) | |
| def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: | |
| return self.net(hidden_states.mean(dim=1)) | |
| class CromaPreTrainedModel(PreTrainedModel): | |
| config_class = CromaConfig | |
| base_model_prefix = "croma" | |
| main_input_name = "optical_pixel_values" | |
| input_modalities = ("image",) | |
| supports_gradient_checkpointing = False | |
| _no_split_modules = ["CromaEncoder", "CromaCrossEncoder"] | |
| def _init_weights(self, module): | |
| if isinstance(module, nn.Linear): | |
| module.weight.data.normal_(mean=0.0, std=self.config.initializer_range) | |
| if module.bias is not None: | |
| module.bias.data.zero_() | |
| elif isinstance(module, nn.LayerNorm): | |
| module.bias.data.zero_() | |
| module.weight.data.fill_(1.0) | |
| class CromaModel(CromaPreTrainedModel): | |
| def __init__(self, config: CromaConfig): | |
| super().__init__(config) | |
| self.modality = config.modality | |
| self.register_buffer( | |
| "attn_bias", | |
| get_2dalibi(config.num_attention_heads, config.num_patches), | |
| persistent=False, | |
| ) | |
| cross_depth = config.num_hidden_layers // 2 | |
| if config.modality in {"sar", "both"}: | |
| self.sar_encoder = CromaViTEncoder(config, depth=cross_depth, in_channels=config.sar_channels) | |
| self.sar_gap_ffn = CromaGapHead(config) | |
| if config.modality in {"optical", "both"}: | |
| self.optical_encoder = CromaViTEncoder( | |
| config, depth=config.num_hidden_layers, in_channels=config.optical_channels | |
| ) | |
| self.optical_gap_ffn = CromaGapHead(config) | |
| if config.modality == "both": | |
| self.cross_encoder = CromaCrossEncoder(config, depth=cross_depth) | |
| self.post_init() | |
| def _resolve_primary_outputs( | |
| self, | |
| sar_hidden_states, | |
| sar_pooler_output, | |
| optical_hidden_states, | |
| optical_pooler_output, | |
| joint_hidden_states, | |
| joint_pooler_output, | |
| ): | |
| if joint_hidden_states is not None: | |
| return joint_hidden_states, joint_pooler_output | |
| if self.modality == "sar" or sar_hidden_states is not None: | |
| return sar_hidden_states, sar_pooler_output | |
| return optical_hidden_states, optical_pooler_output | |
| def forward( | |
| self, | |
| sar_pixel_values: Optional[torch.Tensor] = None, | |
| optical_pixel_values: Optional[torch.Tensor] = None, | |
| return_dict: Optional[bool] = None, | |
| **kwargs: Unpack[TransformersKwargs], | |
| ) -> CromaModelOutput: | |
| if return_dict is None: | |
| return_dict = self.config.use_return_dict | |
| has_sar = sar_pixel_values is not None | |
| has_optical = optical_pixel_values is not None | |
| if self.modality == "both": | |
| if not has_sar and not has_optical: | |
| raise ValueError("Provide at least one of `sar_pixel_values` or `optical_pixel_values`.") | |
| elif self.modality == "sar" and not has_sar: | |
| raise ValueError("Modality is set to 'sar', but `sar_pixel_values` is None.") | |
| elif self.modality == "optical" and not has_optical: | |
| raise ValueError("Modality is set to 'optical', but `optical_pixel_values` is None.") | |
| attn_bias = self.attn_bias | |
| sar_hidden_states = sar_pooler_output = None | |
| optical_hidden_states = optical_pooler_output = None | |
| joint_hidden_states = joint_pooler_output = None | |
| if self.modality in {"sar", "both"} and has_sar: | |
| sar_pixel_values = sar_pixel_values.to(dtype=self.dtype) | |
| attn_bias = attn_bias.to(device=sar_pixel_values.device, dtype=sar_pixel_values.dtype) | |
| sar_hidden_states = self.sar_encoder(sar_pixel_values, attn_bias) | |
| sar_pooler_output = self.sar_gap_ffn(sar_hidden_states) | |
| if self.modality in {"optical", "both"} and has_optical: | |
| optical_pixel_values = optical_pixel_values.to(dtype=self.dtype) | |
| attn_bias = attn_bias.to(device=optical_pixel_values.device, dtype=optical_pixel_values.dtype) | |
| optical_hidden_states = self.optical_encoder(optical_pixel_values, attn_bias) | |
| optical_pooler_output = self.optical_gap_ffn(optical_hidden_states) | |
| if self.modality == "both" and has_sar and has_optical: | |
| joint_hidden_states = self.cross_encoder(sar_hidden_states, optical_hidden_states, attn_bias) | |
| joint_pooler_output = joint_hidden_states.mean(dim=1) | |
| last_hidden_state, pooler_output = self._resolve_primary_outputs( | |
| sar_hidden_states, | |
| sar_pooler_output, | |
| optical_hidden_states, | |
| optical_pooler_output, | |
| joint_hidden_states, | |
| joint_pooler_output, | |
| ) | |
| if not return_dict: | |
| return ( | |
| last_hidden_state, | |
| pooler_output, | |
| sar_hidden_states, | |
| sar_pooler_output, | |
| optical_hidden_states, | |
| optical_pooler_output, | |
| joint_hidden_states, | |
| joint_pooler_output, | |
| ) | |
| return CromaModelOutput( | |
| last_hidden_state=last_hidden_state, | |
| pooler_output=pooler_output, | |
| sar_hidden_states=sar_hidden_states, | |
| sar_pooler_output=sar_pooler_output, | |
| optical_hidden_states=optical_hidden_states, | |
| optical_pooler_output=optical_pooler_output, | |
| joint_hidden_states=joint_hidden_states, | |
| joint_pooler_output=joint_pooler_output, | |
| ) | |
| __all__ = ["CromaConfig", "CromaModel", "CromaModelOutput", "CromaPreTrainedModel"] | |