echo / code /flash-linear-attention /tests /models /test_modeling_mom.py
amonshano's picture
Add Echo-Memory codebase used for this run (CC BY 4.0, JD Echo Team) (part 3)
b66f552 verified
Raw
History Blame Contribute Delete
1.65 kB
import pytest
import torch
from fla.models import MomConfig
from .test_modeling_base import run_test_generation, run_test_model_forward_backward
# ===================================================================================
# Test for Modeling (Forward/Backward Pass)
# ===================================================================================
@pytest.mark.skip(reason="Bug not fixed yet")
@pytest.mark.parametrize(
['L', 'B', 'T', 'H', 'D', 'use_l2warp', 'dtype'],
[
pytest.param(*test, id="L{}-B{}-T{}-H{}-D{}-use_l2warp{}-{}".format(*test))
for test in [
(4, 4, 1024, 4, 64, True, torch.bfloat16),
(4, 4, 1024, 4, 64, False, torch.bfloat16),
(4, 4, 1024, 4, 128, False, torch.bfloat16),
]
],
)
def test_modeling(
L: int,
B: int,
T: int,
H: int,
D: int,
use_l2warp: bool,
dtype: torch.dtype,
):
run_test_model_forward_backward(L, B, T, H, D, MomConfig, use_l2warp=use_l2warp, dtype=dtype)
# ===================================================================================
# Test for Generation
# ===================================================================================
@pytest.mark.parametrize(
['L', 'B', 'T', 'H', 'D', 'dtype'],
[
pytest.param(*test, id="L{}-B{}-T{}-H{}-D{}-{}".format(*test))
for test in [
(2, 4, 2000, 8, 64, torch.float16),
]
],
)
def test_generation(
L: int,
B: int,
T: int,
H: int,
D: int,
dtype: torch.dtype,
):
pytest.skip("Known bugs in mom")
run_test_generation(L, B, T, H, D, MomConfig, dtype)