Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
@@ -1297,7 +1297,7 @@ class CustomQwen2Decoder(nn.Module):
|
|||||||
min_dtype = torch.finfo(dtype).min
|
min_dtype = torch.finfo(dtype).min
|
||||||
|
|
||||||
is_image = token_type_ids == 0 # [B, S]
|
is_image = token_type_ids == 0 # [B, S]
|
||||||
is_text = token_type_ids == 1 # [B, S]
|
is_text = token_type_ids == 1 # [B, S]
|
||||||
|
|
||||||
mask = torch.full(
|
mask = torch.full(
|
||||||
(batch_size, sequence_length, sequence_length),
|
(batch_size, sequence_length, sequence_length),
|
||||||
@@ -1312,8 +1312,8 @@ class CustomQwen2Decoder(nn.Module):
|
|||||||
causal = idx.unsqueeze(0) <= idx.unsqueeze(1) # [S, S]
|
causal = idx.unsqueeze(0) <= idx.unsqueeze(1) # [S, S]
|
||||||
|
|
||||||
text_causal = (
|
text_causal = (
|
||||||
is_text.unsqueeze(2) # [B, S, 1]
|
is_text.unsqueeze(2) # [B, S, 1]
|
||||||
& is_text.unsqueeze(1) # [B, 1, S]
|
& is_text.unsqueeze(1) # [B, 1, S]
|
||||||
& causal.unsqueeze(0) # [1, S, S]
|
& causal.unsqueeze(0) # [1, S, S]
|
||||||
) # [B, S, S]
|
) # [B, S, S]
|
||||||
|
|
||||||
|
|||||||
@@ -43,6 +43,7 @@ _DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
|
|||||||
# Standalone reference implementation (original loop-based code, pre-a475156d)
|
# Standalone reference implementation (original loop-based code, pre-a475156d)
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def _create_custom_4d_mask_reference(
|
def _create_custom_4d_mask_reference(
|
||||||
sequence_length, dtype, device, batch_size, token_type_ids
|
sequence_length, dtype, device, batch_size, token_type_ids
|
||||||
):
|
):
|
||||||
@@ -58,7 +59,7 @@ def _create_custom_4d_mask_reference(
|
|||||||
)
|
)
|
||||||
type_ids = token_type_ids[b]
|
type_ids = token_type_ids[b]
|
||||||
image_positions = (type_ids == 0).nonzero(as_tuple=True)[0]
|
image_positions = (type_ids == 0).nonzero(as_tuple=True)[0]
|
||||||
text_positions = (type_ids == 1).nonzero(as_tuple=True)[0]
|
text_positions = (type_ids == 1).nonzero(as_tuple=True)[0]
|
||||||
|
|
||||||
if len(image_positions) > 0:
|
if len(image_positions) > 0:
|
||||||
mask[image_positions[:, None], image_positions] = 0.0
|
mask[image_positions[:, None], image_positions] = 0.0
|
||||||
@@ -79,13 +80,14 @@ def _create_custom_4d_mask_reference(
|
|||||||
# python/sglang/srt/models/deepseek_ocr.py)
|
# python/sglang/srt/models/deepseek_ocr.py)
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def _create_custom_4d_mask_new(
|
def _create_custom_4d_mask_new(
|
||||||
sequence_length, dtype, device, batch_size, token_type_ids
|
sequence_length, dtype, device, batch_size, token_type_ids
|
||||||
):
|
):
|
||||||
min_dtype = torch.finfo(dtype).min
|
min_dtype = torch.finfo(dtype).min
|
||||||
|
|
||||||
is_image = token_type_ids == 0 # [B, S]
|
is_image = token_type_ids == 0 # [B, S]
|
||||||
is_text = token_type_ids == 1 # [B, S]
|
is_text = token_type_ids == 1 # [B, S]
|
||||||
|
|
||||||
mask = torch.full(
|
mask = torch.full(
|
||||||
(batch_size, sequence_length, sequence_length),
|
(batch_size, sequence_length, sequence_length),
|
||||||
@@ -100,8 +102,8 @@ def _create_custom_4d_mask_new(
|
|||||||
causal = idx.unsqueeze(0) <= idx.unsqueeze(1) # [S, S]
|
causal = idx.unsqueeze(0) <= idx.unsqueeze(1) # [S, S]
|
||||||
|
|
||||||
text_causal = (
|
text_causal = (
|
||||||
is_text.unsqueeze(2) # [B, S, 1]
|
is_text.unsqueeze(2) # [B, S, 1]
|
||||||
& is_text.unsqueeze(1) # [B, 1, S]
|
& is_text.unsqueeze(1) # [B, 1, S]
|
||||||
& causal.unsqueeze(0) # [1, S, S]
|
& causal.unsqueeze(0) # [1, S, S]
|
||||||
) # [B, S, S]
|
) # [B, S, S]
|
||||||
|
|
||||||
@@ -117,6 +119,7 @@ def _create_custom_4d_mask_new(
|
|||||||
# Helpers
|
# Helpers
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def _make_token_type_ids(batch_size, seq_len, image_fraction, device):
|
def _make_token_type_ids(batch_size, seq_len, image_fraction, device):
|
||||||
"""First `image_fraction` tokens per sequence are image (0), rest are text (1).
|
"""First `image_fraction` tokens per sequence are image (0), rest are text (1).
|
||||||
|
|
||||||
@@ -139,7 +142,7 @@ def _make_random_token_type_ids(batch_size, seq_len, device, seed=42):
|
|||||||
def _bench_cuda_events(fn, n, **kwargs):
|
def _bench_cuda_events(fn, n, **kwargs):
|
||||||
"""Time `fn` on CUDA using cuda events (excludes H2D launch overhead)."""
|
"""Time `fn` on CUDA using cuda events (excludes H2D launch overhead)."""
|
||||||
start = torch.cuda.Event(enable_timing=True)
|
start = torch.cuda.Event(enable_timing=True)
|
||||||
end = torch.cuda.Event(enable_timing=True)
|
end = torch.cuda.Event(enable_timing=True)
|
||||||
# warmup
|
# warmup
|
||||||
for _ in range(5):
|
for _ in range(5):
|
||||||
fn(**kwargs)
|
fn(**kwargs)
|
||||||
@@ -149,7 +152,7 @@ def _bench_cuda_events(fn, n, **kwargs):
|
|||||||
fn(**kwargs)
|
fn(**kwargs)
|
||||||
end.record()
|
end.record()
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
return start.elapsed_time(end) / 1e3 / n # seconds per iteration
|
return start.elapsed_time(end) / 1e3 / n # seconds per iteration
|
||||||
|
|
||||||
|
|
||||||
def _bench_wall(fn, n, **kwargs):
|
def _bench_wall(fn, n, **kwargs):
|
||||||
@@ -172,13 +175,14 @@ def _bench(fn, run_device, n=50, **kwargs):
|
|||||||
# Accuracy tests
|
# Accuracy tests
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
class TestAccuracy(unittest.TestCase):
|
class TestAccuracy(unittest.TestCase):
|
||||||
"""Verify new implementation produces identical masks to the reference."""
|
"""Verify new implementation produces identical masks to the reference."""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
cls.device = _DEVICE
|
cls.device = _DEVICE
|
||||||
cls.dtype = torch.float32
|
cls.dtype = torch.float32
|
||||||
|
|
||||||
def _check(self, batch_size, seq_len, token_type_ids):
|
def _check(self, batch_size, seq_len, token_type_ids):
|
||||||
ref = _create_custom_4d_mask_reference(
|
ref = _create_custom_4d_mask_reference(
|
||||||
@@ -241,10 +245,10 @@ class TestAccuracy(unittest.TestCase):
|
|||||||
def test_batch_heterogeneous(self):
|
def test_batch_heterogeneous(self):
|
||||||
"""Different image/text ratios per batch item."""
|
"""Different image/text ratios per batch item."""
|
||||||
ids = torch.ones(4, 64, dtype=torch.long, device=self.device)
|
ids = torch.ones(4, 64, dtype=torch.long, device=self.device)
|
||||||
ids[0, :10] = 0
|
ids[0, :10] = 0
|
||||||
ids[1, :32] = 0
|
ids[1, :32] = 0
|
||||||
ids[2, :63] = 0
|
ids[2, :63] = 0
|
||||||
ids[3, :] = 1
|
ids[3, :] = 1
|
||||||
self._check(4, 64, ids)
|
self._check(4, 64, ids)
|
||||||
|
|
||||||
# --- output shape ---
|
# --- output shape ---
|
||||||
@@ -273,44 +277,52 @@ class TestAccuracy(unittest.TestCase):
|
|||||||
def test_causal_text_ordering(self):
|
def test_causal_text_ordering(self):
|
||||||
"""Text token i must NOT attend to text token j > i."""
|
"""Text token i must NOT attend to text token j > i."""
|
||||||
B, S = 1, 8
|
B, S = 1, 8
|
||||||
ids = torch.ones(B, S, dtype=torch.long, device=self.device)
|
ids = torch.ones(B, S, dtype=torch.long, device=self.device)
|
||||||
out = _create_custom_4d_mask_new(S, self.dtype, self.device, B, ids)
|
out = _create_custom_4d_mask_new(S, self.dtype, self.device, B, ids)
|
||||||
min_val = torch.finfo(self.dtype).min
|
min_val = torch.finfo(self.dtype).min
|
||||||
mask2d = out.cpu()[0, 0]
|
mask2d = out.cpu()[0, 0]
|
||||||
for q in range(S):
|
for q in range(S):
|
||||||
for k in range(S):
|
for k in range(S):
|
||||||
if k <= q:
|
if k <= q:
|
||||||
self.assertEqual(mask2d[q, k].item(), 0.0,
|
self.assertEqual(
|
||||||
f"text[{q}] should attend to text[{k}]")
|
mask2d[q, k].item(),
|
||||||
|
0.0,
|
||||||
|
f"text[{q}] should attend to text[{k}]",
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
self.assertEqual(mask2d[q, k].item(), min_val,
|
self.assertEqual(
|
||||||
f"text[{q}] should NOT attend to text[{k}]")
|
mask2d[q, k].item(),
|
||||||
|
min_val,
|
||||||
|
f"text[{q}] should NOT attend to text[{k}]",
|
||||||
|
)
|
||||||
|
|
||||||
def test_image_full_attention(self):
|
def test_image_full_attention(self):
|
||||||
"""Image tokens must attend to all other image tokens (bidirectional)."""
|
"""Image tokens must attend to all other image tokens (bidirectional)."""
|
||||||
B, S = 1, 12
|
B, S = 1, 12
|
||||||
n_img = 6
|
n_img = 6
|
||||||
ids = torch.ones(B, S, dtype=torch.long, device=self.device)
|
ids = torch.ones(B, S, dtype=torch.long, device=self.device)
|
||||||
ids[:, :n_img] = 0
|
ids[:, :n_img] = 0
|
||||||
out = _create_custom_4d_mask_new(S, self.dtype, self.device, B, ids)
|
out = _create_custom_4d_mask_new(S, self.dtype, self.device, B, ids)
|
||||||
mask2d = out.cpu()[0, 0]
|
mask2d = out.cpu()[0, 0]
|
||||||
for q in range(n_img):
|
for q in range(n_img):
|
||||||
for k in range(n_img):
|
for k in range(n_img):
|
||||||
self.assertEqual(mask2d[q, k].item(), 0.0,
|
self.assertEqual(
|
||||||
f"image[{q}] should attend to image[{k}]")
|
mask2d[q, k].item(), 0.0, f"image[{q}] should attend to image[{k}]"
|
||||||
|
)
|
||||||
|
|
||||||
def test_text_attends_to_image(self):
|
def test_text_attends_to_image(self):
|
||||||
"""Every text token must attend to every image token."""
|
"""Every text token must attend to every image token."""
|
||||||
B, S = 1, 12
|
B, S = 1, 12
|
||||||
n_img = 4
|
n_img = 4
|
||||||
ids = torch.ones(B, S, dtype=torch.long, device=self.device)
|
ids = torch.ones(B, S, dtype=torch.long, device=self.device)
|
||||||
ids[:, :n_img] = 0
|
ids[:, :n_img] = 0
|
||||||
out = _create_custom_4d_mask_new(S, self.dtype, self.device, B, ids)
|
out = _create_custom_4d_mask_new(S, self.dtype, self.device, B, ids)
|
||||||
mask2d = out.cpu()[0, 0]
|
mask2d = out.cpu()[0, 0]
|
||||||
for q in range(n_img, S):
|
for q in range(n_img, S):
|
||||||
for k in range(n_img):
|
for k in range(n_img):
|
||||||
self.assertEqual(mask2d[q, k].item(), 0.0,
|
self.assertEqual(
|
||||||
f"text[{q}] should attend to image[{k}]")
|
mask2d[q, k].item(), 0.0, f"text[{q}] should attend to image[{k}]"
|
||||||
|
)
|
||||||
|
|
||||||
# --- dtype coverage ---
|
# --- dtype coverage ---
|
||||||
|
|
||||||
@@ -333,14 +345,14 @@ class TestAccuracy(unittest.TestCase):
|
|||||||
|
|
||||||
BENCHMARK_CASES = [
|
BENCHMARK_CASES = [
|
||||||
# (batch_size, seq_len, image_fraction)
|
# (batch_size, seq_len, image_fraction)
|
||||||
(1, 256, 0.5),
|
(1, 256, 0.5),
|
||||||
(4, 512, 0.5),
|
(4, 512, 0.5),
|
||||||
(8, 1024, 0.5),
|
(8, 1024, 0.5),
|
||||||
(16, 2048, 0.5),
|
(16, 2048, 0.5),
|
||||||
(4, 4096, 0.75),
|
(4, 4096, 0.75),
|
||||||
]
|
]
|
||||||
BENCH_ITERS = 50
|
BENCH_ITERS = 50
|
||||||
SPEEDUP_FLOOR = 1.0 # new must be at least as fast as reference
|
SPEEDUP_FLOOR = 1.0 # new must be at least as fast as reference
|
||||||
|
|
||||||
|
|
||||||
class TestPerformance(unittest.TestCase):
|
class TestPerformance(unittest.TestCase):
|
||||||
@@ -349,11 +361,12 @@ class TestPerformance(unittest.TestCase):
|
|||||||
@classmethod
|
@classmethod
|
||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
cls.device = _DEVICE
|
cls.device = _DEVICE
|
||||||
cls.dtype = torch.float32
|
cls.dtype = torch.float32
|
||||||
|
|
||||||
def _run_case(self, batch_size, seq_len, image_fraction):
|
def _run_case(self, batch_size, seq_len, image_fraction):
|
||||||
ids = _make_token_type_ids(batch_size, seq_len, image_fraction,
|
ids = _make_token_type_ids(
|
||||||
device=self.device)
|
batch_size, seq_len, image_fraction, device=self.device
|
||||||
|
)
|
||||||
kwargs = dict(
|
kwargs = dict(
|
||||||
sequence_length=seq_len,
|
sequence_length=seq_len,
|
||||||
dtype=self.dtype,
|
dtype=self.dtype,
|
||||||
@@ -361,10 +374,15 @@ class TestPerformance(unittest.TestCase):
|
|||||||
batch_size=batch_size,
|
batch_size=batch_size,
|
||||||
token_type_ids=ids,
|
token_type_ids=ids,
|
||||||
)
|
)
|
||||||
t_ref = _bench(_create_custom_4d_mask_reference, run_device=self.device,
|
t_ref = _bench(
|
||||||
n=BENCH_ITERS, **kwargs)
|
_create_custom_4d_mask_reference,
|
||||||
t_new = _bench(_create_custom_4d_mask_new, run_device=self.device,
|
run_device=self.device,
|
||||||
n=BENCH_ITERS, **kwargs)
|
n=BENCH_ITERS,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
t_new = _bench(
|
||||||
|
_create_custom_4d_mask_new, run_device=self.device, n=BENCH_ITERS, **kwargs
|
||||||
|
)
|
||||||
speedup = t_ref / t_new
|
speedup = t_ref / t_new
|
||||||
dev_tag = "CUDA" if "cuda" in str(self.device) else "CPU"
|
dev_tag = "CUDA" if "cuda" in str(self.device) else "CPU"
|
||||||
print(
|
print(
|
||||||
@@ -406,6 +424,7 @@ class TestPerformance(unittest.TestCase):
|
|||||||
# PyTorch profiler (optional – triggered by --profile or PROFILE_TRACES=1)
|
# PyTorch profiler (optional – triggered by --profile or PROFILE_TRACES=1)
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def run_profiler_traces(output_dir: str = "./pt_traces", device: str = _DEVICE):
|
def run_profiler_traces(output_dir: str = "./pt_traces", device: str = _DEVICE):
|
||||||
"""
|
"""
|
||||||
Capture Chrome-trace JSON files for both implementations.
|
Capture Chrome-trace JSON files for both implementations.
|
||||||
@@ -440,7 +459,7 @@ def run_profiler_traces(output_dir: str = "./pt_traces", device: str = _DEVICE):
|
|||||||
|
|
||||||
for label, fn in [
|
for label, fn in [
|
||||||
("reference", _create_custom_4d_mask_reference),
|
("reference", _create_custom_4d_mask_reference),
|
||||||
("new", _create_custom_4d_mask_new),
|
("new", _create_custom_4d_mask_new),
|
||||||
]:
|
]:
|
||||||
trace_path = os.path.join(
|
trace_path = os.path.join(
|
||||||
output_dir,
|
output_dir,
|
||||||
@@ -495,7 +514,7 @@ if __name__ == "__main__":
|
|||||||
"--device",
|
"--device",
|
||||||
default=None,
|
default=None,
|
||||||
help="Device to run on: 'cuda', 'cuda:0', 'cpu', etc. "
|
help="Device to run on: 'cuda', 'cuda:0', 'cpu', etc. "
|
||||||
"Defaults to CUDA if available, otherwise CPU.",
|
"Defaults to CUDA if available, otherwise CPU.",
|
||||||
)
|
)
|
||||||
args, remaining = parser.parse_known_args()
|
args, remaining = parser.parse_known_args()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user