diff --git a/src/modalities/models/gpt2/gpt2_model.py b/src/modalities/models/gpt2/gpt2_model.py index 993221e2c..3d5e5bacd 100644 --- a/src/modalities/models/gpt2/gpt2_model.py +++ b/src/modalities/models/gpt2/gpt2_model.py @@ -393,7 +393,7 @@ def _update_cos_sin_tables(self, x): ): self._seq_len_cached = seq_len t = torch.arange(x.shape[self.seq_length_dim], device=x.device, dtype=torch.float32) - freqs = torch.einsum("i,j->ij", t, self.inv_freq.to(x.dtype)) + freqs = torch.einsum("i,j->ij", t, self.inv_freq.float()) emb = torch.cat((freqs, freqs), dim=-1).to( x.device ) # here, we combine the two matrices (not zipping them). diff --git a/tests/test_rotary_qkv_transform.py b/tests/test_rotary_qkv_transform.py index b44868e4b..81fcacaf5 100644 --- a/tests/test_rotary_qkv_transform.py +++ b/tests/test_rotary_qkv_transform.py @@ -44,6 +44,29 @@ def test_rotary_transform(): assert torch.equal(comp_rot_expected, comp_rot) +def test_rotary_transform_computes_frequencies_in_fp32_for_bf16_inputs(monkeypatch): + operand_and_result_dtypes = [] + original_einsum = torch.einsum + + def recording_einsum(equation, *operands): + result = original_einsum(equation, *operands) + operand_and_result_dtypes.append((*[operand.dtype for operand in operands], result.dtype)) + return result + + monkeypatch.setattr(torch, "einsum", recording_einsum) + + rotary_transform = RotaryTransform(n_embd=128, n_head=2) + q = torch.randn(1, 2, 16, 64, dtype=torch.bfloat16) + k = torch.randn_like(q) + v = torch.randn_like(q) + + q_rot, k_rot, _ = rotary_transform(q=q, k=k, v=v) + + assert operand_and_result_dtypes == [(torch.float32, torch.float32, torch.float32)] + assert q_rot.dtype == torch.bfloat16 + assert k_rot.dtype == torch.bfloat16 + + def _apply_rotary(x: torch.Tensor, cos_cached: torch.Tensor, sin_cached: torch.Tensor) -> torch.Tensor: cos_local = cos_cached[:, :, : x.shape[-2], :] sin_local = sin_cached[:, :, : x.shape[-2], :] @@ -61,7 +84,7 @@ def _assert_yarn_outputs_match_reference( seq_length: int, ) -> None: t = torch.arange(seq_length, device=q.device, dtype=torch.float32) - freqs = torch.einsum("i,j->ij", t, rotary_transform.inv_freq.to(q.dtype)) + freqs = torch.einsum("i,j->ij", t, rotary_transform.inv_freq.float()) emb = torch.cat((freqs, freqs), dim=-1) cos = (emb.cos() * rotary_transform.attention_scaling)[None, None, :, :].to(q.dtype) sin = (emb.sin() * rotary_transform.attention_scaling)[None, None, :, :].to(q.dtype)