[diffusion] Fix test_model_fast_paths import after sana_ln_modulate rename (#36175)
Co-authored-by: Alex Nails <alex.nails@radixark.ai>
This commit is contained in:
@@ -99,7 +99,7 @@ from sglang.multimodal_gen.runtime.models.dits.sana import (
|
|||||||
_eager_ln_modulate as _sana_eager_ln_modulate,
|
_eager_ln_modulate as _sana_eager_ln_modulate,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.models.dits.sana import (
|
from sglang.multimodal_gen.runtime.models.dits.sana import (
|
||||||
_sana_ln_modulate,
|
sana_ln_modulate,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.models.vaes import flux2_vae_cuda_opt as vae_opt
|
from sglang.multimodal_gen.runtime.models.vaes import flux2_vae_cuda_opt as vae_opt
|
||||||
from sglang.multimodal_gen.runtime.models.vaes.autoencoder import AutoencoderKL
|
from sglang.multimodal_gen.runtime.models.vaes.autoencoder import AutoencoderKL
|
||||||
@@ -361,7 +361,7 @@ def test_sana_fused_ln_modulate_is_bit_exact(shape, nmod, transposed):
|
|||||||
shift, scale = emb.chunk(nmod, dim=1)[0], emb.chunk(nmod, dim=1)[-1]
|
shift, scale = emb.chunk(nmod, dim=1)[0], emb.chunk(nmod, dim=1)[-1]
|
||||||
# default-stream eager serving must stay on the untouched eager chain
|
# default-stream eager serving must stay on the untouched eager chain
|
||||||
n_sigs = len(sana._SANA_LN_MOD.verified_sigs)
|
n_sigs = len(sana._SANA_LN_MOD.verified_sigs)
|
||||||
_sana_ln_modulate(norm, x, scale, shift)
|
sana_ln_modulate(norm, x, scale, shift)
|
||||||
assert len(sana._SANA_LN_MOD.verified_sigs) == n_sigs
|
assert len(sana._SANA_LN_MOD.verified_sigs) == n_sigs
|
||||||
# The fusion engages on non-default streams (the BCG warmup/capture path).
|
# The fusion engages on non-default streams (the BCG warmup/capture path).
|
||||||
# x/scale/shift were filled on the default stream, so the side stream must
|
# x/scale/shift were filled on the default stream, so the side stream must
|
||||||
@@ -373,9 +373,9 @@ def test_sana_fused_ln_modulate_is_bit_exact(shape, nmod, transposed):
|
|||||||
side = torch.cuda.Stream()
|
side = torch.cuda.Stream()
|
||||||
side.wait_stream(torch.cuda.current_stream())
|
side.wait_stream(torch.cuda.current_stream())
|
||||||
with torch.cuda.stream(side):
|
with torch.cuda.stream(side):
|
||||||
out = _sana_ln_modulate(norm, x, scale, shift)
|
out = sana_ln_modulate(norm, x, scale, shift)
|
||||||
assert len(sana._SANA_LN_MOD.verified_sigs) == n_sigs + 1 # verified
|
assert len(sana._SANA_LN_MOD.verified_sigs) == n_sigs + 1 # verified
|
||||||
out2 = _sana_ln_modulate(norm, x, scale, shift) # verified-sig lane
|
out2 = sana_ln_modulate(norm, x, scale, shift) # verified-sig lane
|
||||||
torch.cuda.current_stream().wait_stream(side)
|
torch.cuda.current_stream().wait_stream(side)
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
assert torch.equal(out, _sana_eager_ln_modulate(norm, x, scale, shift))
|
assert torch.equal(out, _sana_eager_ln_modulate(norm, x, scale, shift))
|
||||||
|
|||||||
Reference in New Issue
Block a user