MoCha: revert some RoPE changes

This commit is contained in:
kijai
2025-10-21 20:12:40 +03:00
parent 5c1f64197e
commit 089329ef2e
+49 -57
View File
@@ -17,73 +17,65 @@ def rope_params_mocha(max_seq_len, dim, theta=10000, L_test=25, k=0, start=0):
@torch.autocast(device_type=mm.get_autocast_device(mm.get_torch_device()), enabled=False)
def rope_apply_mocha(x, grid_sizes, freqs):
batch_size, _, n, c_doubled = x.shape
c = c_doubled // 2
n, c = x.size(2), x.size(3) // 2
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
split_size_1 = c - 2 * (c // 3)
split_size_2 = c // 3
split_size_3 = c // 3
freqs_0, freqs_1, freqs_2 = freqs.split([split_size_1, split_size_2, split_size_3], dim=1)
f_tensor = grid_sizes[0, 0]
h_tensor = grid_sizes[0, 1]
w_tensor = grid_sizes[0, 2]
batch_size = grid_sizes.size(0)
outputs = []
seq_len_tensor = f_tensor * h_tensor * w_tensor
sf_tensor = (f_tensor - 2) // 2
for i in range(batch_size):
f, h, w = grid_sizes[i].tolist()
seq_len = f * h * w
sf_range = torch.arange(1, sf_tensor + 1, device=freqs[0].device)
h_range = torch.arange(1, h_tensor + 1, device=freqs[1].device)
w_range = torch.arange(1, w_tensor + 1, device=freqs[2].device)
repeat_freqs = torch.cat([
freqs[0][sf_range].view(sf_tensor, 1, 1, -1).expand(sf_tensor, h_tensor, w_tensor, -1),
freqs[1][h_range].view(1, h_tensor, 1, -1).expand(sf_tensor, h_tensor, w_tensor, -1),
freqs[2][w_range].view(1, 1, w_tensor, -1).expand(sf_tensor, h_tensor, w_tensor, -1)
], dim=-1)
x_i = torch.view_as_complex(x[i, :seq_len].reshape(seq_len, n, -1, 2))
mask_freqs = torch.cat([
freqs[0][1:2].view(1, 1, 1, -1).expand(1, h_tensor, w_tensor, -1),
freqs[1][h_range].view(1, h_tensor, 1, -1).expand(1, h_tensor, w_tensor, -1),
freqs[2][w_range].view(1, 1, w_tensor, -1).expand(1, h_tensor, w_tensor, -1)
], dim=-1)
sf = (f - 2) // 2
img_freqs = torch.cat([
freqs[0][0:1].view(1, 1, 1, -1).expand(1, h_tensor, w_tensor, -1),
freqs[1][h_range].view(1, h_tensor, 1, -1).expand(1, h_tensor, w_tensor, -1),
freqs[2][w_range].view(1, 1, w_tensor, -1).expand(1, h_tensor, w_tensor, -1)
], dim=-1)
freq0_start, freq0_end = 1, 1 + sf
freq1_start, freq1_end = 1, 1 + h
freq2_start, freq2_end = 1, 1 + w
repeat_freqs = torch.cat([
freqs_0[freq0_start:freq0_end].view(sf, 1, 1, -1).expand(sf, h, w, -1),
freqs_1[freq1_start:freq1_end].view(1, h, 1, -1).expand(sf, h, w, -1),
freqs_2[freq2_start:freq2_end].view(1, 1, w, -1).expand(sf, h, w, -1)
], dim=-1)
condition = (f_tensor == 2 * sf_tensor + 2)
bias_h_range = torch.arange(h_tensor + 1, 2 * h_tensor + 1, device=freqs[1].device)
bias_w_range = torch.arange(w_tensor + 1, 2 * w_tensor + 1, device=freqs[2].device)
bias_freqs = torch.cat([
freqs[0][0:1].view(1, 1, 1, -1).expand(1, h_tensor, w_tensor, -1),
freqs[1][bias_h_range].view(1, h_tensor, 1, -1).expand(1, h_tensor, w_tensor, -1),
freqs[2][bias_w_range].view(1, 1, w_tensor, -1).expand(1, h_tensor, w_tensor, -1)
], dim=-1)
freqs_without_bias = torch.cat([repeat_freqs, repeat_freqs, mask_freqs, img_freqs], dim=0)
freqs_with_bias = torch.cat([repeat_freqs, repeat_freqs, mask_freqs, img_freqs, bias_freqs], dim=0)
pad_size = freqs_with_bias.size(0) - freqs_without_bias.size(0)
padding = torch.zeros(
pad_size, h_tensor, w_tensor, freqs_without_bias.size(-1),
dtype=freqs_without_bias.dtype, device=freqs_without_bias.device
)
freqs_without_bias = torch.cat([freqs_without_bias, padding], dim=0)
freqs_i = torch.where(condition.unsqueeze(0).unsqueeze(0).unsqueeze(0), freqs_without_bias, freqs_with_bias)
freqs_i = freqs_i.reshape(seq_len_tensor, 1, -1).to(x.device)
mask_freqs = torch.cat([
freqs_0[1:2].view(1, 1, 1, -1).expand(1, h, w, -1),
freqs_1[freq1_start:freq1_end].view(1, h, 1, -1).expand(1, h, w, -1),
freqs_2[freq2_start:freq2_end].view(1, 1, w, -1).expand(1, h, w, -1)
], dim=-1)
x_seq = x[:, :seq_len_tensor].to(torch.float64).reshape(batch_size, seq_len_tensor, n, -1, 2)
x_complex = torch.view_as_complex(x_seq)
img_freqs = torch.cat([
freqs_0[0:1].view(1, 1, 1, -1).expand(1, h, w, -1),
freqs_1[freq1_start:freq1_end].view(1, h, 1, -1).expand(1, h, w, -1),
freqs_2[freq2_start:freq2_end].view(1, 1, w, -1).expand(1, h, w, -1)
], dim=-1)
x_rotated = torch.view_as_real(x_complex * freqs_i.unsqueeze(0)).flatten(3)
if f == 2 * sf + 2:
freqs_i = torch.cat([repeat_freqs, repeat_freqs, mask_freqs, img_freqs], dim=0)
else:
bias_freqs = torch.cat([
freqs_0[0:1].view(1, 1, 1, -1).expand(1, h, w, -1),
freqs_1[(h+1):(2*h+1)].view(1, h, 1, -1).expand(1, h, w, -1),
freqs_2[(w+1):(2*w+1)].view(1, 1, w, -1).expand(1, h, w, -1)
], dim=-1)
freqs_i = torch.cat([repeat_freqs, repeat_freqs, mask_freqs, img_freqs, bias_freqs], dim=0)
freqs_i = freqs_i.reshape(f * h * w, 1, -1).to(x.device)
# Apply rotary embedding
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
x_i = torch.cat([x_i, x[i, seq_len:]], dim=0)
outputs.append(x_i)
x_remaining = x[:, seq_len_tensor:]
output = torch.cat([x_rotated, x_remaining], dim=1)
return output.to(x.dtype)
return torch.stack(outputs).to(x.dtype)
device = mm.get_torch_device()