Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
317 changes: 317 additions & 0 deletions tests/pytorch/test_torch_compile.py
Original file line number Diff line number Diff line change
Expand Up @@ -2720,3 +2720,320 @@ def test_te_ops_forward_kwargs_compile():
for gain in (3.0, 5.0):
_check_ops(compiled, model, x, dy, {"gain": gain})
_assert_custom_ops(graphs[-1:], "_affineop", present=False)


@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available")
@pytest.mark.parametrize("compile_mode", ["default", "reduce-overhead"])
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32])
@pytest.mark.parametrize("equal_splits", [False, True])
def test_te_split_parameters_compile(compile_mode, dtype, equal_splits):
"""Every split remains an autograd input across storage and parameter changes."""
import copy

torch._dynamo.reset()
counters.clear()
splits = ("q", "k", "v") if equal_splits else dict(q=64, k=32, v=32)
model = te.Linear(64, 192 if equal_splits else 128, parameters_split=splits, params_dtype=dtype)
model.k_weight.requires_grad_(False)
model.v_bias.requires_grad_(False)
reference = copy.deepcopy(model)
compiled = torch.compile(model, fullgraph=True, mode=compile_mode)
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
ref_optimizer = torch.optim.SGD(reference.parameters(), lr=0.01)
for mutation in ("none", "optimizer", "data", "parameter", "assign", "dtype"):
for current in (model, reference):
if mutation == "data":
current.k_weight.data = current.k_weight.detach().clone() + 0.01
current.k_bias.data = current.k_bias.detach().clone() + 0.01
elif mutation == "parameter":
current.q_weight = torch.nn.Parameter(current.q_weight.detach().clone() + 0.01)
elif mutation == "assign":
state = {name: value.clone() for name, value in current.state_dict().items()}
current.load_state_dict(state, assign=True)
elif mutation == "dtype":
current.to(torch.float64).to(dtype)
if mutation == "optimizer":
optimizer.step()
ref_optimizer.step()
for _ in range(3):
torch.compiler.cudagraph_mark_step_begin()
inp = torch.randn(16, 64, device="cuda", dtype=dtype, requires_grad=True)
ref_inp = inp.detach().clone().requires_grad_()
model.zero_grad(set_to_none=True)
reference.zero_grad(set_to_none=True)
out = compiled(inp)
grad = torch.randn_like(out)
out.backward(grad)
actual = [out.detach().clone(), inp.grad.clone()]
actual.extend(None if p.grad is None else p.grad.clone() for p in model.parameters())
ref_out = reference(ref_inp)
ref_out.backward(grad)
expected = [ref_out, ref_inp.grad, *(p.grad for p in reference.parameters())]
for result, target in zip(actual, expected):
if target is None:
assert result is None
else:
torch.testing.assert_close(result, target, **dtype_tols(dtype))
if compile_mode == "reduce-overhead":
assert not counters["inductor"]["cudagraph_skips"]


@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available")
@pytest.mark.parametrize("fp8_recipe", [None, *_all_recipes], ids=recipe_id)
@pytest.mark.parametrize("bias_mode", ["fused", "none", "returned"])
def test_te_split_parameters_recipes(fp8_recipe, bias_mode):
autocast_recipe = fp8_recipe or recipe.Float8CurrentScaling()
options = dict(bias=bias_mode != "none", return_bias=bias_mode == "returned")
model = te.Linear(
128,
256,
parameters_split=dict(q=128, k=64, v=64),
params_dtype=torch.bfloat16,
**options,
)

def fn(inp):
with te.autocast(enabled=fp8_recipe is not None, recipe=autocast_recipe):
out = model(inp)
if bias_mode == "returned":
out, bias = out
out = out + bias
return out

torch._dynamo.reset()
compiled = torch.compile(fn, fullgraph=True)
inp = torch.randn(128, 128, device="cuda", dtype=torch.bfloat16, requires_grad=True)
grad = torch.randn(128, 256, device="cuda", dtype=torch.bfloat16)
expected = fn(inp)
expected.backward(grad)
expected_grads = [t.grad.clone() for t in (inp, *model.parameters()) if t.requires_grad]
inp.grad = None
model.zero_grad(set_to_none=True)
actual = compiled(inp)
actual.backward(grad)
actual_grads = [t.grad for t in (inp, *model.parameters()) if t.requires_grad]
torch.testing.assert_close(actual, expected, rtol=0, atol=0)
torch.testing.assert_close(actual_grads, expected_grads, rtol=0, atol=0)


@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available")
@pytest.mark.parametrize("training", [False, True])
def test_te_split_parameters_no_cat(training):
"""Adjacent parameter storage does not launch a concatenation in either pass."""
model = te.Linear(64, 128, parameters_split=dict(q=64, k=32, v=32), params_dtype=torch.bfloat16)
compiled = torch.compile(model, fullgraph=True)
inp = torch.randn(16, 64, device="cuda", dtype=torch.bfloat16, requires_grad=training)

def run():
with torch.set_grad_enabled(training):
out = compiled(inp)
if training:
out.sum().backward()
model.zero_grad(set_to_none=True)
inp.grad = None

for _ in range(3):
run()
with torch.profiler.profile(activities=[torch.profiler.ProfilerActivity.CPU]) as prof:
run()
assert "aten::cat" not in {event.key for event in prof.key_averages()}
model.k_weight.data = model.k_weight.detach().clone()
with torch.profiler.profile(activities=[torch.profiler.ProfilerActivity.CPU]) as prof:
run()
assert "aten::cat" in {event.key for event in prof.key_averages()}


@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available")
def test_te_split_parameters_saved_versions():
model = te.Linear(64, 128, parameters_split=dict(q=64, k=32, v=32), params_dtype=torch.bfloat16)
compiled = torch.compile(model, fullgraph=True, backend="aot_eager")
inp = torch.randn(16, 64, device="cuda", dtype=torch.bfloat16, requires_grad=True)
out = compiled(inp)
with torch.no_grad():
model.v_weight.add_(0.01)
with pytest.raises(RuntimeError, match="modified by an inplace operation"):
out.sum().backward()


@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available")
def test_te_split_parameters_saved_hooks():
import copy

model = te.Linear(64, 128, parameters_split=dict(q=64, k=32, v=32), params_dtype=torch.bfloat16)
reference = copy.deepcopy(model)
compiled = torch.compile(model, fullgraph=True, backend="aot_eager")
inp = torch.randn(16, 64, device="cuda", dtype=torch.bfloat16, requires_grad=True)
ref_inp = inp.detach().clone().requires_grad_()
with torch.autograd.graph.saved_tensors_hooks(lambda t: t.clone(), lambda t: t):
out = compiled(inp)
out.sum().backward()
expected = reference(ref_inp)
expected.sum().backward()
torch.testing.assert_close(out, expected, rtol=0, atol=0)
torch.testing.assert_close(inp.grad, ref_inp.grad, rtol=0, atol=0)
for actual, target in zip(model.parameters(), reference.parameters()):
torch.testing.assert_close(actual.grad, target.grad, rtol=0, atol=0)


def test_te_split_parameters_fallback_check(monkeypatch):
model = te.Linear(64, 128, parameters_split=dict(q=64, k=32, v=32))
inp = torch.empty(16, 64, device="cuda")

def unexpected_materialization(*args, **kwargs):
raise AssertionError("Fallback eligibility must not concatenate parameters")

monkeypatch.setattr(model, "_get_weight_and_bias_tensors", unexpected_materialization)
assert model._compile_eager_fallback_reason(inp, None, False, False, True, False) is None
assert (
model._compile_eager_fallback_reason(inp, None, True, False, True, False)
== "differentiable fp8_output=True"
)


@pytest.mark.parametrize("defer", [False, True])
def test_concat_input_single_tensor(defer):
from transformer_engine.pytorch.module._common import concat_input

tensor = torch.nn.Parameter(torch.randn(4, 8, device="cuda"))
assert concat_input([tensor], defer=defer) is tensor


@pytest.mark.parametrize("disjoint", [False, True])
def test_concat_input_eager_gradients(disjoint):
from transformer_engine.pytorch.module._common import concat_input

parts = [torch.nn.Parameter(t) for t in torch.randn(8, 8, device="cuda").split((3, 5))]
if disjoint:
parts[1] = torch.nn.Parameter(parts[1].detach().clone())
result = concat_input(parts)
expected = torch.cat(parts)
grad = torch.randn_like(result)
torch.testing.assert_close(result, expected, rtol=0, atol=0)
torch.testing.assert_close(
torch.autograd.grad(result, parts, grad),
torch.autograd.grad(expected, parts, grad),
rtol=0,
atol=0,
)


@pytest.mark.parametrize(
"layout", ["adjacent", "offset", "disjoint", "strided", "singleton", "negative", "conjugate"]
)
def test_parameter_parts_storage(layout):
from transformer_engine.pytorch.dynamo.parameter_parts import ParameterParts

storage = torch.arange(128, device="cuda").view(16, 8)
if layout == "offset":
storage = storage[2:10]
elif layout == "strided":
storage = storage[:, ::2]
elif layout == "singleton":
storage = storage[:, :1]
elif layout == "negative":
storage = torch._neg_view(storage)
elif layout == "conjugate":
storage = torch.complex(storage.float(), storage.float()).conj()
parts = list(storage.split([1, storage.shape[0] - 3, 2]))
if layout == "disjoint":
parts[1] = parts[1].clone()
result = ParameterParts(parts).materialize()
torch.testing.assert_close(result, torch.cat(parts), rtol=0, atol=0)
if layout in ("adjacent", "offset"):
assert result.data_ptr() == parts[0].data_ptr()
else:
assert result.untyped_storage().data_ptr() != parts[0].untyped_storage().data_ptr()


@pytest.mark.parametrize("fake", [False, True], ids=["eager", "fake"])
@pytest.mark.parametrize("mixed_dtype", [False, True])
def test_parameter_parts_spec(fake, mixed_dtype):
from transformer_engine.pytorch.dynamo.parameter_parts import ParameterParts

with FakeTensorMode() if fake else contextlib.nullcontext():
parts = [
torch.empty(3, 8, device="cuda", dtype=torch.bfloat16),
torch.empty(
5,
8,
device="cuda",
dtype=torch.float32 if mixed_dtype else torch.bfloat16,
requires_grad=True,
),
]
expected = torch.cat(parts)
spec = ParameterParts(parts).to_spec()
assert spec.shape == tuple(expected.shape)
assert spec.dtype == expected.dtype
assert spec.device == expected.device
assert spec.requires_grad == expected.requires_grad
assert not spec.is_quantized


@pytest.mark.parametrize("internal", [False, True])
def test_parameter_parts_rejects_quantized_parts(internal):
from transformer_engine.pytorch.dynamo.parameter_parts import ParameterParts
from transformer_engine.pytorch.module._common import concat_input

quantizer = _current_scaling()
quantizer.internal = internal
tensor = TensorSpec(
shape=(4, 8), dtype=torch.bfloat16, quantizer=quantizer, device=torch.device("cpu")
).create_tensor()
with pytest.raises(TypeError, match="non-quantized tensors"):
ParameterParts([tensor, tensor])
assert concat_input([tensor], defer=True) is tensor


@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available")
@pytest.mark.parametrize("grad_target", ["input", "weight", "bias"])
def test_te_split_parameters_grad_targets(grad_target):
model = te.Linear(64, 128, parameters_split=dict(q=64, k=32, v=32), params_dtype=torch.bfloat16)
for parameter in model.parameters():
parameter.requires_grad_(False)
if grad_target != "input":
getattr(model, f"v_{grad_target}").requires_grad_(True)
inp = torch.randn(16, 64, device="cuda", dtype=torch.bfloat16)
inp.requires_grad_(grad_target == "input")
compiled = torch.compile(model, fullgraph=True)
out = compiled(inp)
grad = torch.randn_like(out)
out.backward(grad)
if grad_target == "bias":
torch.testing.assert_close(model.v_bias.grad, grad[:, -32:].sum(0), rtol=0, atol=0)
actual = [out.detach().clone()]
tensors = [inp, *model.parameters()]
actual.extend(None if t.grad is None else t.grad.clone() for t in tensors)
for tensor in tensors:
tensor.grad = None
expected = model(inp)
expected.backward(grad)
for result, target in zip(actual, [expected, *(t.grad for t in tensors)]):
if target is None:
assert result is None
else:
torch.testing.assert_close(result, target, rtol=0, atol=0)


@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available")
@pytest.mark.parametrize("bf16_part", ["q", "k"])
def test_te_split_parameters_mixed_dtype_autocast(bf16_part):
model = te.Linear(64, 192, parameters_split=("q", "k", "v"), params_dtype=torch.float32)
name = f"{bf16_part}_weight"
setattr(model, name, torch.nn.Parameter(getattr(model, name).to(torch.bfloat16)))
inp = torch.randn(16, 64, device="cuda", dtype=torch.float32, requires_grad=True)
grad = torch.randn(16, 192, device="cuda", dtype=torch.bfloat16)
torch._dynamo.reset()
compiled = torch.compile(model, fullgraph=True)
with torch.autocast("cuda", dtype=torch.bfloat16):
expected = model(inp)
expected.backward(grad)
expected_grads = [tensor.grad.clone() for tensor in (inp, *model.parameters())]
inp.grad = None
model.zero_grad(set_to_none=True)
with torch.autocast("cuda", dtype=torch.bfloat16):
actual = compiled(inp)
actual.backward(grad)
actual_grads = [tensor.grad for tensor in (inp, *model.parameters())]
torch.testing.assert_close(actual, expected, rtol=0, atol=0)
torch.testing.assert_close(actual_grads, expected_grads, rtol=0, atol=0)
Loading
Loading