of3gs-demo / src /model /encoder /pe_utils.py
richardchencccc's picture
Add OF3GS ZeroGPU demo
f737f60 verified
Raw
History Blame Contribute Delete
5.15 kB
import torch
import math
import numpy as np
def get_1d_sincos_pos_emb_from_grid(embed_dim, pos, device="cpu"):
"""
Generate 1D sinusoidal positional embeddings from grid positions.
Args:
embed_dim (int): The embedding dimension (must be even).
pos (torch.Tensor): The grid positions (e.g., [0, 1, 2, ..., v-1]).
Shape: [b * gh * gw] or [batch_size, sequence_length].
device (str): Device for the output tensor.
Returns:
torch.Tensor: Sinusoidal positional embeddings.
Shape: [len(pos), embed_dim]
"""
assert embed_dim % 2 == 0, "Embedding dimension must be even for sine and cosine."
# Convert positions to float
pos = pos.float()
# Compute the sinusoidal frequencies
dim = torch.arange(embed_dim // 2, dtype=torch.float32, device=device) # [0, 1, ..., embed_dim // 2 - 1]
freq = 1.0 / (10000 ** (dim / (embed_dim // 2))) # Scale frequencies logarithmically
# Calculate sine and cosine embeddings
pos_emb_sin = torch.sin(pos[:, None] * freq) # Shape: [len(pos), embed_dim // 2]
pos_emb_cos = torch.cos(pos[:, None] * freq) # Shape: [len(pos), embed_dim // 2]
# Concatenate sine and cosine along the last dimension
pos_emb = torch.cat([pos_emb_sin, pos_emb_cos], dim=-1) # Shape: [len(pos), embed_dim]
return pos_emb
def get_2d_sincos_pos_embed(embed_dim, grid_size, device="cpu"):
"""
Generate 2D sine-cosine positional embeddings with separate grid height and width.
Args:
embed_dim (int): The embedding dimension.
grid_size (tuple): Tuple specifying the grid height and width (grid_h, grid_w).
cls_token (bool): Whether to include a [CLS] token embedding.
device (str): The device to place the embeddings on.
Returns:
torch.Tensor: Positional embeddings of shape [grid_h*grid_w, embed_dim]
"""
grid_h, grid_w = grid_size # Unpack grid dimensions
# Create the grid for height and width
grid_h = torch.arange(grid_h, dtype=torch.float32, device=device)
grid_w = torch.arange(grid_w, dtype=torch.float32, device=device)
grid = torch.meshgrid(grid_w, grid_h, indexing="ij") # w goes first
grid = torch.stack(grid, dim=0) # Shape: [2, grid_h, grid_w]
# Reshape grid to [2, 1, grid_h, grid_w]
grid = grid.view(2, 1, grid.size(1), grid.size(2))
# Get the positional embeddings from the grid
pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid, device=device) # Shape: [grid_h*grid_w, embed_dim]
return pos_embed
def get_2d_sincos_pos_embed_from_grid(embed_dim, grid, device="cpu"):
"""
Generate 2D sine-cosine positional embeddings from a grid with separate height and width.
Args:
embed_dim (int): The embedding dimension.
grid (torch.Tensor): The grid of shape [2, 1, grid_h, grid_w].
device (str): The device to place the embeddings on.
Returns:
torch.Tensor: Positional embeddings of shape [grid_h*grid_w, embed_dim].
"""
assert embed_dim % 2 == 0, "Embedding dimension must be even."
# Use half of the dimensions for each grid dimension
grid_h = grid[0].view(-1) # Flatten grid height dimension: [H*W]
grid_w = grid[1].view(-1) # Flatten grid width dimension: [H*W]
# Generate 1D sine-cosine embeddings for grid_h and grid_w
emb_h = get_1d_sincos_pos_emb_from_grid(embed_dim // 2, grid_h, device=device) # Shape: [H*W, D/2]
emb_w = get_1d_sincos_pos_emb_from_grid(embed_dim // 2, grid_w, device=device) # Shape: [H*W, D/2]
# Concatenate along the last dimension
pos_embed = torch.cat([emb_h, emb_w], dim=-1) # Shape: [H*W, D]
return pos_embed
def rope(positions: torch.Tensor, d: int, device='cpu') -> torch.Tensor:
"""
Given a batch of positions in [0,1], compute RoPE-style
sine-cosine embeddings in dimension d (must be even).
positions: (B, N) tensor of float positions in [0,1].
d: int, dimension of the embedding (should be even).
Returns:
embeddings: (B, N, d) tensor of float embeddings.
"""
# positions shape: [B, N]
B, N = positions.shape
half_d = d // 2
# Expand positions to shape [B, N, 1]
positions_3d = positions.unsqueeze(-1) # [B, N, 1]
# Prepare index and frequency tensors
# idx => [1, 1, half_d]
idx = torch.arange(half_d, device=device).view(1, 1, -1)
# freqs => [1, 1, half_d], broadcast to [B, N, half_d]
freqs = torch.pow(10000.0, -2.0 * idx / d)
# angle => [B, N, half_d]
angle = positions_3d.to(device) * freqs
# Compute sine and cosine => each [B, N, half_d]
sin_part = angle.sin()
cos_part = angle.cos()
# Interleave sine and cosine along the last dimension => [B, N, d]
embeddings = torch.empty(B, N, d, device=device, dtype=positions.dtype)
embeddings[..., 0::2] = sin_part
embeddings[..., 1::2] = cos_part
return embeddings