File size: 5,149 Bytes
f737f60
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
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