AFR-DFV-v2 / dinov3 /layers /ffn_layers.py
Addax-Data-Science's picture
Upload 162 files
d9bb75c verified
Raw
History Blame Contribute Delete
2.59 kB
# Copyright (c) Meta Platforms, Inc. and affiliates.
#
# This software may be used and distributed in accordance with
# the terms of the DINOv3 License Agreement.
from typing import Callable, List, Optional
import torch.nn.functional as F
from torch import Tensor, nn
from dinov3.utils import cat_keep_shapes, uncat_with_shapes
class ListForwardMixin(object):
def forward(self, x: Tensor):
raise NotImplementedError
def forward_list(self, x_list: List[Tensor]) -> List[Tensor]:
x_flat, shapes, num_tokens = cat_keep_shapes(x_list)
x_flat = self.forward(x_flat)
return uncat_with_shapes(x_flat, shapes, num_tokens)
class Mlp(nn.Module, ListForwardMixin):
def __init__(
self,
in_features: int,
hidden_features: Optional[int] = None,
out_features: Optional[int] = None,
act_layer: Callable[..., nn.Module] = nn.GELU,
drop: float = 0.0,
bias: bool = True,
device=None,
) -> None:
super().__init__()
out_features = out_features or in_features
hidden_features = hidden_features or in_features
self.fc1 = nn.Linear(in_features, hidden_features, bias=bias, device=device)
self.act = act_layer()
self.fc2 = nn.Linear(hidden_features, out_features, bias=bias, device=device)
self.drop = nn.Dropout(drop)
def forward(self, x: Tensor) -> Tensor:
x = self.fc1(x)
x = self.act(x)
x = self.drop(x)
x = self.fc2(x)
x = self.drop(x)
return x
class SwiGLUFFN(nn.Module, ListForwardMixin):
def __init__(
self,
in_features: int,
hidden_features: Optional[int] = None,
out_features: Optional[int] = None,
act_layer: Optional[Callable[..., nn.Module]] = None,
drop: float = 0.0,
bias: bool = True,
align_to: int = 8,
device=None,
) -> None:
super().__init__()
out_features = out_features or in_features
hidden_features = hidden_features or in_features
d = int(hidden_features * 2 / 3)
swiglu_hidden_features = d + (-d % align_to)
self.w1 = nn.Linear(in_features, swiglu_hidden_features, bias=bias, device=device)
self.w2 = nn.Linear(in_features, swiglu_hidden_features, bias=bias, device=device)
self.w3 = nn.Linear(swiglu_hidden_features, out_features, bias=bias, device=device)
def forward(self, x: Tensor) -> Tensor:
x1 = self.w1(x)
x2 = self.w2(x)
hidden = F.silu(x1) * x2
return self.w3(hidden)