first commit
This commit is contained in:
35
models/positional_encoding.py
Normal file
35
models/positional_encoding.py
Normal file
@@ -0,0 +1,35 @@
|
||||
"""Standard 2D sine-cosine positional encoding, normalized by grid size.
|
||||
|
||||
Normalizing by grid size (rather than absolute pixel index, as the paper's Eq. 14
|
||||
ambiguously implies) keeps PE values independent of input resolution -- one of the
|
||||
camera/resolution-agnostic design choices in the project plan.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class PositionEmbedding2D(nn.Module):
|
||||
def __init__(self, dim: int, temperature: float = 10000.0):
|
||||
super().__init__()
|
||||
assert dim % 4 == 0, "dim must be divisible by 4 for 2D sine-cosine PE"
|
||||
self.dim = dim
|
||||
self.temperature = temperature
|
||||
|
||||
def forward(self, h: int, w: int, device: torch.device) -> torch.Tensor:
|
||||
"""Returns (h*w, dim) position embedding."""
|
||||
num_pos_feats = self.dim // 2 # half the channels for x, half for y
|
||||
y_embed = torch.linspace(0, 1, h, device=device).unsqueeze(1).repeat(1, w) # (h, w)
|
||||
x_embed = torch.linspace(0, 1, w, device=device).unsqueeze(0).repeat(h, 1) # (h, w)
|
||||
|
||||
dim_t = torch.arange(num_pos_feats, device=device, dtype=torch.float32)
|
||||
dim_t = self.temperature ** (2 * (dim_t // 2) / num_pos_feats)
|
||||
|
||||
pos_x = x_embed[..., None] / dim_t
|
||||
pos_y = y_embed[..., None] / dim_t
|
||||
pos_x = torch.stack((pos_x[..., 0::2].sin(), pos_x[..., 1::2].cos()), dim=-1).flatten(-2)
|
||||
pos_y = torch.stack((pos_y[..., 0::2].sin(), pos_y[..., 1::2].cos()), dim=-1).flatten(-2)
|
||||
|
||||
pos = torch.cat((pos_y, pos_x), dim=-1) # (h, w, dim)
|
||||
return pos.flatten(0, 1) # (h*w, dim)
|
||||
Reference in New Issue
Block a user