amonshano's picture
Add Echo-Memory codebase used for this run (CC BY 4.0, JD Echo Team) (part 4)
c335050 verified
Raw
History Blame Contribute Delete
2.59 kB
import pytest
import torch
from fla.modules.grpo import fused_grpo_loss, grpo_loss_torch
from fla.utils import assert_close, device, device_torch_lib, is_nvidia_hopper
@pytest.mark.parametrize("B", [2])
@pytest.mark.parametrize("T", [16, 1024, 4096])
@pytest.mark.parametrize("V", [32000, 65536, 131072])
@pytest.mark.parametrize("dtype", [torch.bfloat16])
@pytest.mark.parametrize("inplace", [True, False])
@pytest.mark.parametrize("repeat", [100])
def test_fused_grpos(B: int, T: int, V: int, dtype: torch.dtype, inplace: bool, repeat: int):
device_torch_lib.manual_seed(42)
for i in range(repeat):
if not is_nvidia_hopper and T == 4096:
pytest.skip("Skip test for T=4096 on Intel Alchemist")
def get_random_ref_log_probs(logits, input_ids):
with torch.inference_mode():
logits = logits[:, :-1]
per_token_logps = []
for logits_row, input_ids_row in zip(logits, input_ids[:, -logits.size(1):], strict=False):
log_probs = torch.randn_like(logits_row).log_softmax(dim=-1)
token_log_prob = torch.gather(log_probs, dim=1, index=input_ids_row.unsqueeze(1)).squeeze(1)
per_token_logps.append(token_log_prob)
device_torch_lib.empty_cache()
return torch.stack(per_token_logps)
logits = torch.randn(B, T + 1, V, device=device, dtype=dtype)
logits.requires_grad_(True)
advantages = torch.randn(B, device=device, dtype=torch.float32)
input_ids = torch.randint(0, V-1, (B, T + 64), device=device)
ref_logp = get_random_ref_log_probs(logits, input_ids)
beta = 0.04
completion_mask = torch.ones(B, T, dtype=torch.int32, device=device)
completion_mask[::2, T//3: T//2] = 0
save_kl = True
gold_logits = logits.detach().clone().float()
gold_logits.requires_grad_(True)
gold_ref_logp = ref_logp.clone().float()
device_torch_lib.empty_cache()
y1 = fused_grpo_loss(logits, ref_logp, input_ids, advantages, beta, completion_mask, save_kl=save_kl, inplace=inplace)
y2 = grpo_loss_torch(gold_logits, gold_ref_logp, input_ids, advantages, beta, completion_mask, save_kl)
if save_kl:
y1, kl2 = y1
y2, kl3 = y2
assert (kl2-kl3).abs().max() < 1e-3
dy = torch.randn_like(y1) * 10
y1.backward(dy)
y2.backward(dy.float())
assert (y1-y2).abs().max() < 1e-3
assert_close(" dlogits", gold_logits.grad, logits.grad, 3e-3)