BiliSakura's picture
Upload 12 files
eb1c9a0 verified
Raw
History Blame Contribute Delete
13.4 kB
# 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"]