diff --git a/python/sglang/srt/layers/attention/xpu_backend.py b/python/sglang/srt/layers/attention/xpu_backend.py index d56459962..900d0fbd7 100644 --- a/python/sglang/srt/layers/attention/xpu_backend.py +++ b/python/sglang/srt/layers/attention/xpu_backend.py @@ -110,6 +110,21 @@ class XPUAttentionBackend(AttentionBackend): 1 if get_exec().deterministic.enable_deterministic_inference else 0 ) self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder + if self.is_encoder_decoder: + from sglang.srt.model_executor.cuda_graph_config import ( + cuda_graph_fully_disabled, + ) + + # Encoder-decoder cross-/self-attention below uses a dynamic-shape + # varlen KV gather (page_size=1 semantics) that cannot be captured. + # XPU disables CUDA graph by default, so this holds; the guard fails + # loudly if a future XPU graph path is force-enabled instead of + # silently mis-indexing through the paged graph-metadata path. + assert cuda_graph_fully_disabled(), ( + "Encoder-decoder models (e.g. Whisper) on the intel_xpu attention " + "backend require CUDA graph disabled (off by default on XPU); the " + "graph decode path cannot run the varlen KV gather." + ) def init_forward_metadata(self, forward_batch: ForwardBatch): """Initialize forward metadata hence all layers in the forward pass can reuse it.""" @@ -361,10 +376,6 @@ class XPUAttentionBackend(AttentionBackend): # Encoder metadata for cross attention if forward_batch.encoder_lens is not None: - assert forward_batch.encoder_lens.numel() == 1, ( - "Only encoder size 1 is supported for now" - ) - metadata.encoder_lens_int32 = forward_batch.encoder_lens.to(torch.int32) metadata.encoder_cu_seqlens_k = torch.nn.functional.pad( torch.cumsum(metadata.encoder_lens_int32, dim=0, dtype=torch.int32), @@ -375,12 +386,18 @@ class XPUAttentionBackend(AttentionBackend): forward_batch.req_pool_indices, : metadata.encoder_max_seq_len_k ] - # Currently only support forward_batch.encoder_lens.numel() == 1 + # Decoder self-attn KV: per-request token-granular slice starting at + # each request's own encoder offset encoder_lens[i], not a single max. + text_max = metadata.max_seq_len_k + arange_text = torch.arange( + text_max, device=forward_batch.req_pool_indices.device + ) + text_col = forward_batch.encoder_lens.long().unsqueeze( + 1 + ) + arange_text.unsqueeze(0) + text_row = forward_batch.req_pool_indices.unsqueeze(1).expand(-1, text_max) metadata.page_table = self.req_to_token_pool.req_to_token[ - forward_batch.req_pool_indices, - metadata.encoder_max_seq_len_k : ( - metadata.encoder_max_seq_len_k + metadata.max_seq_len_k - ), + text_row, text_col ] # Translate full-pool indices to SWA-pool indices for hybrid models @@ -415,8 +432,10 @@ class XPUAttentionBackend(AttentionBackend): workspace_size, device=self.device, dtype=torch.uint8 ) - # Convert the page table to a strided format which is needed by FA3 API - if self.page_size > 1: + # Convert the page table to a strided format which is needed by FA3 API. + # Encoder-decoder page_table holds token-slot indices for the varlen + # kernel (page_size=1 semantics), so it must not be page-strided. + if self.page_size > 1 and forward_batch.encoder_lens is None: self.strided_indices = torch.arange( 0, metadata.page_table.shape[1], self.page_size, device=self.device ) @@ -572,6 +591,22 @@ class XPUAttentionBackend(AttentionBackend): value_cache = value_cache.view( -1, self.page_size, layer.tp_v_head_num, layer.head_dim ) + if self.is_encoder_decoder and forward_batch.encoder_lens is not None: + page_table, cache_seqlens, causal = self._encoder_decoder_page_table( + layer, metadata + ) + o = self._forward_attn_flat_page_table( + q=q, + key_cache=key_cache, + value_cache=value_cache, + layer=layer, + page_table=page_table, + cache_seqlens=cache_seqlens, + cu_seqlens_q=metadata.cu_seqlens_q, + max_seqlen_q=metadata.max_seq_len_q, + causal=causal, + ) + return o.view(-1, layer.tp_q_head_num * layer.v_head_dim) if layer.is_cross_attention: page_table = metadata.encoder_page_table cache_seqlens = metadata.encoder_lens_int32 @@ -765,6 +800,64 @@ class XPUAttentionBackend(AttentionBackend): out = o.view(-1, layer.tp_q_head_num * layer.v_head_dim) return out + @staticmethod + def _encoder_decoder_page_table(layer, metadata): + """Pick (page_table, cache_seqlens, causal) for an encoder-decoder layer: + cross-attention reads the encoder KV region (non-causal), decoder + self-attention reads the decoder KV region (causal).""" + if layer.is_cross_attention: + return metadata.encoder_page_table, metadata.encoder_lens_int32, False + return metadata.page_table, metadata.cache_seqlens_int32, True + + def _forward_attn_flat_page_table( + self, + *, + q, + key_cache, + value_cache, + layer, + page_table, + cache_seqlens, + cu_seqlens_q, + max_seqlen_q, + causal, + ): + """MHA on XPU via flash_attn_with_kvcache with a page_size=1 (flat + token-slot) page table. sgl-kernel-xpu PR #454 detects a page_size==1 + k_cache + page_table and gathers + runs varlen internally, so the backend + calls it like the FA (CUDA) path. Eager-only. A request with + cache_seqlens==0 attends to no keys and the kernel returns NaN for its + rows, so those rows are zeroed (an all-empty batch skips the launch). + """ + q_rows = q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim) + # If the maximum cache_seqlens is 0, there are no keys to attend to. + if int(cache_seqlens.max().item()) == 0: + return q_rows.new_zeros( + (q_rows.shape[0], q_rows.shape[1], value_cache.shape[-1]) + ) + # page_size=1 view of the KV pool: PR #454 detects if page size dim is 1 + # i.e. k_cache.shape[1] == 1 and routes flash_attn_with_kvcache to varlen gather. + k_cache = key_cache.reshape(-1, 1, layer.tp_k_head_num, layer.head_dim) + v_cache = value_cache.reshape(-1, 1, layer.tp_v_head_num, layer.head_dim) + out = flash_attn_with_kvcache( + q=q_rows, + k_cache=k_cache, + v_cache=v_cache, + page_table=page_table, + cache_seqlens=cache_seqlens, + cu_seqlens_q=cu_seqlens_q, + max_seqlen_q=max_seqlen_q, + softmax_scale=layer.scaling, + causal=causal, + softcap=layer.logit_cap, + ) + # Mixed batch: requests with cache_seqlens==0 attend to no keys and come + # back as NaN, so zero their query rows (mapped via cu_seqlens_q). + if int(cache_seqlens.min().item()) == 0: + seg = cu_seqlens_q[1:] - cu_seqlens_q[:-1] + out[(cache_seqlens == 0).repeat_interleave(seg)] = 0 + return out + def forward_decode( self, q: torch.Tensor, @@ -870,6 +963,23 @@ class XPUAttentionBackend(AttentionBackend): -1, self.page_size, layer.tp_v_head_num, layer.head_dim ) + if self.is_encoder_decoder and forward_batch.encoder_lens is not None: + page_table, cache_seqlens, causal = self._encoder_decoder_page_table( + layer, metadata + ) + o = self._forward_attn_flat_page_table( + q=q, + key_cache=key_cache, + value_cache=value_cache, + layer=layer, + page_table=page_table, + cache_seqlens=cache_seqlens, + cu_seqlens_q=metadata.cu_seqlens_q, + max_seqlen_q=1, + causal=causal, + ) + return o.view(-1, layer.tp_q_head_num * layer.v_head_dim) + if layer.is_cross_attention: # Always use non-chunked logic for cross-attention o = flash_attn_with_kvcache( diff --git a/test/registered/unit/layers/attention/test_encoder_decoder_varlen_gather.py b/test/registered/unit/layers/attention/test_encoder_decoder_varlen_gather.py new file mode 100644 index 000000000..afa15db5f --- /dev/null +++ b/test/registered/unit/layers/attention/test_encoder_decoder_varlen_gather.py @@ -0,0 +1,219 @@ +import sys +import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import torch + +from sglang.test.ci.ci_register import register_cpu_ci + +with patch.dict( + sys.modules, + { + module: MagicMock() + for module in ( + "sgl_kernel", + "sgl_kernel.flash_attn", + "sgl_kernel.quantization", + "sgl_kernel.scalar_type", + ) + }, +): + from sglang.srt.layers.attention import xpu_backend + from sglang.srt.layers.attention.xpu_backend import XPUAttentionBackend + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class TestEncoderDecoderForward(unittest.TestCase): + """Encoder-decoder attention on XPU calls flash_attn_with_kvcache with a + page_size=1 view; sgl-kernel-xpu PR #454 gathers the token-slot page_table and + runs varlen inside the kernel. These guard the backend's own responsibilities + -- the cross-vs-self dispatch, the page_size=1 view, the empty-KV zero guard, + and the per-request encoder-offset metadata. Real-kernel numerical correctness + is covered on device in test/registered/xpu/test_xpu_encoder_decoder_varlen.py. + """ + + HQ, HK, D, TOTAL_SLOTS = 4, 2, 8, 40 + + def setUp(self): + torch.manual_seed(0) + self.backend = XPUAttentionBackend.__new__(XPUAttentionBackend) + self.k_flat = torch.randn(self.TOTAL_SLOTS, self.HK, self.D) + self.v_flat = torch.randn(self.TOTAL_SLOTS, self.HK, self.D) + + def _layer(self, is_cross): + return SimpleNamespace( + is_cross_attention=is_cross, + tp_q_head_num=self.HQ, + tp_k_head_num=self.HK, + tp_v_head_num=self.HK, + head_dim=self.D, + scaling=0.5, + logit_cap=0.0, + ) + + def test_dispatch_and_forward_cross_vs_self(self): + # The caller picks (page_table, cache_seqlens, causal) via + # _encoder_decoder_page_table -- cross-attn -> encoder_page_table + + # encoder_lens_int32 + causal=False; self-attn -> page_table + + # cache_seqlens_int32 + causal=True -- then hands them to the generic + # _forward_attn_flat_page_table, which must forward them unchanged with a + # page_size=1 k_cache (shape[1]==1) so PR #454 routes to the varlen gather. + enc_pt = torch.arange(5, dtype=torch.int32).unsqueeze(0) + dec_pt = (torch.arange(4, dtype=torch.int32) + 10).unsqueeze(0) + metadata = SimpleNamespace( + encoder_page_table=enc_pt, + encoder_lens_int32=torch.tensor([5], dtype=torch.int32), + page_table=dec_pt, + cache_seqlens_int32=torch.tensor([4], dtype=torch.int32), + ) + key_cache = self.k_flat.view(-1, 1, self.HK, self.D) + value_cache = self.v_flat.view(-1, 1, self.HK, self.D) + q = torch.randn(1, self.HQ * self.D) + cu_seqlens_q = torch.tensor([0, 1], dtype=torch.int32) + + for is_cross, exp_pt, exp_seqlens, exp_causal in ( + (True, enc_pt, metadata.encoder_lens_int32, False), + (False, dec_pt, metadata.cache_seqlens_int32, True), + ): + layer = self._layer(is_cross) + page_table, cache_seqlens, causal = ( + self.backend._encoder_decoder_page_table(layer, metadata) + ) + self.assertTrue(torch.equal(page_table, exp_pt)) + self.assertTrue(torch.equal(cache_seqlens, exp_seqlens)) + self.assertEqual(causal, exp_causal) + + captured = {} + + def fake_kvcache(*_, **kw): + captured.update(kw) + return kw["q"].new_zeros( + (kw["q"].shape[0], kw["q"].shape[1], kw["v_cache"].shape[-1]) + ) + + with patch.object(xpu_backend, "flash_attn_with_kvcache", fake_kvcache): + self.backend._forward_attn_flat_page_table( + q=q, + key_cache=key_cache, + value_cache=value_cache, + layer=layer, + page_table=page_table, + cache_seqlens=cache_seqlens, + cu_seqlens_q=cu_seqlens_q, + max_seqlen_q=1, + causal=causal, + ) + self.assertTrue(torch.equal(captured["page_table"], exp_pt)) + self.assertTrue(torch.equal(captured["cache_seqlens"], exp_seqlens)) + self.assertEqual(captured["causal"], exp_causal) + self.assertEqual(captured["k_cache"].shape[1], 1) # page_size=1 view + self.assertEqual(captured["max_seqlen_q"], 1) + + def test_all_empty_returns_zeros_without_kernel(self): + # Whisper text-only warmup: all cache_seqlens == 0. PR #454's page_size=1 + # path returns NaN for an empty KV, so the backend must short-circuit to + # zeros and never launch the kernel. + key_cache = self.k_flat.view(-1, 1, self.HK, self.D) + value_cache = self.v_flat.view(-1, 1, self.HK, self.D) + q = torch.randn(1, self.HQ * self.D) + sentinel = MagicMock(side_effect=AssertionError("kernel must not run")) + with patch.object(xpu_backend, "flash_attn_with_kvcache", sentinel): + out = self.backend._forward_attn_flat_page_table( + q=q, + key_cache=key_cache, + value_cache=value_cache, + layer=self._layer(True), + page_table=torch.zeros(1, 0, dtype=torch.int32), + cache_seqlens=torch.zeros(1, dtype=torch.int32), + cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32), + max_seqlen_q=1, + causal=False, + ) + sentinel.assert_not_called() + self.assertTrue(torch.equal(out, torch.zeros(1, self.HQ, self.D))) + + def test_mixed_empty_zeros_only_empty_request_rows(self): + # Mixed batch: request 0 has cache_seqlens==0 (no keys), request 1 has + # keys. PR #454 returns NaN for the empty request's rows, so the backend + # must zero exactly those rows and leave the rest untouched. Unequal query + # counts (2 and 3) exercise the cu_seqlens_q -> per-request row mapping. + key_cache = self.k_flat.view(-1, 1, self.HK, self.D) + value_cache = self.v_flat.view(-1, 1, self.HK, self.D) + q = torch.randn(5, self.HQ * self.D) + + def fake_kvcache(*_, **kw): + # All-ones (never-NaN) sentinel so zeroed rows are distinguishable. + return kw["q"].new_ones( + (kw["q"].shape[0], kw["q"].shape[1], kw["v_cache"].shape[-1]) + ) + + with patch.object(xpu_backend, "flash_attn_with_kvcache", fake_kvcache): + out = self.backend._forward_attn_flat_page_table( + q=q, + key_cache=key_cache, + value_cache=value_cache, + layer=self._layer(True), + page_table=torch.zeros(2, 4, dtype=torch.int32), + cache_seqlens=torch.tensor([0, 4], dtype=torch.int32), + cu_seqlens_q=torch.tensor([0, 2, 5], dtype=torch.int32), + max_seqlen_q=3, + causal=False, + ) + self.assertTrue(torch.equal(out[:2], torch.zeros(2, self.HQ, self.D))) + self.assertTrue(torch.equal(out[2:], torch.ones(3, self.HQ, self.D))) + + def test_init_forward_metadata_per_request_encoder_offset(self): + # Guards the encoder_lens.numel()==1 removal: with UNEQUAL encoder lengths, + # init_forward_metadata must place each request's decoder self-attn page + # table at ITS OWN encoder_lens[i] offset (not a single batch max), and + # slice encoder KV per request. It must also skip the //page_size stride + # for enc-dec (token-slot indices feed the varlen kernel). Fails on the old + # scalar-max-offset slice (row0 would start at col 5, not 3). + from sglang.srt.model_executor.forward_batch_info import ForwardMode + + backend = XPUAttentionBackend.__new__(XPUAttentionBackend) + backend.page_size = 128 # >1: also exercises the enc-dec stride-skip + backend.is_encoder_decoder = True + backend.use_mla = False + backend.use_sliding_window_kv_pool = False + backend.attention_chunk_size = None + backend.topk = 0 + # req_to_token[i, j] = 100*i + j, so gathered values reveal (row, col). + req_to_token = torch.arange(16).unsqueeze(0) + torch.tensor([[0], [100]]) + backend.req_to_token_pool = SimpleNamespace(req_to_token=req_to_token) + + fb = SimpleNamespace( + forward_mode=ForwardMode.DECODE, + seq_lens=torch.tensor([2, 4], dtype=torch.int64), # decoder lengths + seq_lens_cpu=torch.tensor([2, 4]), + batch_size=2, + req_pool_indices=torch.tensor([0, 1]), + encoder_lens=torch.tensor([3, 5], dtype=torch.int64), # UNEQUAL + spec_info=None, + out_cache_loc=None, + ) + backend.init_forward_metadata(fb) + md = backend.forward_metadata + + # Encoder KV: columns [0 : max_enc=5] of each request's row; per-request + # lengths + segment boundaries captured for the kernel's internal gather. + self.assertEqual( + md.encoder_page_table.tolist(), + [[0, 1, 2, 3, 4], [100, 101, 102, 103, 104]], + ) + self.assertEqual(md.encoder_lens_int32.tolist(), [3, 5]) + self.assertEqual(md.encoder_cu_seqlens_k.tolist(), [0, 3, 8]) + + # Decoder self-attn KV (text_max = max(seq_lens) = 4 columns each): request 0 + # starts at col 3 (its encoder_len), request 1 at col 5 (its encoder_len). + # Token-granular (not //128), proving the stride transform was skipped. + self.assertEqual( + md.page_table.tolist(), + [[3, 4, 5, 6], [105, 106, 107, 108]], + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/xpu/test_xpu_encoder_decoder_varlen.py b/test/registered/xpu/test_xpu_encoder_decoder_varlen.py new file mode 100644 index 000000000..403e49340 --- /dev/null +++ b/test/registered/xpu/test_xpu_encoder_decoder_varlen.py @@ -0,0 +1,160 @@ +"""Real-device XPU tests for the encoder-decoder attention path. + +The backend calls flash_attn_with_kvcache with a page_size=1 view; sgl-kernel-xpu +PR #454 detects that and gathers + runs varlen inside the kernel. This runs on an +actual XPU and guards what a mocked CPU test cannot: _forward_attn_flat_page_table +plus the real kernel produce correct attention for a scattered (non-page-aligned) +token-slot layout, for both cross-attn (non-causal) and decoder self-attn (causal). +""" + +import unittest +from types import SimpleNamespace + +import torch + +from sglang.srt.layers.attention.xpu_backend import XPUAttentionBackend +from sglang.test.ci.ci_register import register_xpu_ci +from sglang.test.test_utils import CustomTestCase + +register_xpu_ci(est_time=15, suite="stage-b-test-1-gpu-xpu") + + +def _sdpa_ref( + *, q, k_flat, v_flat, page_table, cache_seqlens, cu_seqlens_q, scale, causal +): + """Per-request SDPA oracle: gather each request's valid slots straight from + the page_table and attend. fp32 math for a stable bf16 comparison.""" + outs = [] + for i in range(page_table.shape[0]): + qs, qe = int(cu_seqlens_q[i]), int(cu_seqlens_q[i + 1]) + sl = int(cache_seqlens[i]) + q_i = q[qs:qe].float() + slots = page_table[i, :sl].long() + k_i = k_flat.index_select(0, slots).float() + v_i = v_flat.index_select(0, slots).float() + scores = torch.einsum("qhd,khd->hqk", q_i, k_i) * scale + lq, lk = qe - qs, sl + if causal and lq > 0 and lk > 0: + row = torch.arange(lq, device=q.device).unsqueeze(1) + col = torch.arange(lk, device=q.device).unsqueeze(0) + keep = col <= (lk - lq + row) + scores = scores.masked_fill(~keep.unsqueeze(0), float("-inf")) + probs = scores.softmax(dim=-1) + outs.append(torch.einsum("hqk,khd->qhd", probs, v_i)) + return torch.cat(outs, dim=0) + + +@unittest.skipUnless( + hasattr(torch, "xpu") and torch.xpu.is_available(), "requires an Intel XPU" +) +class TestXPUEncoderDecoderVarlen(CustomTestCase): + # Whisper-large-v3 is MHA (num_kv_heads == num_heads), so use MHA here to keep + # the reference exact; head_dim=64 satisfies the kernel's alignment. + H, D, TOTAL_SLOTS = 8, 64, 64 + + def setUp(self): + torch.manual_seed(0) + self.dev = torch.device("xpu") + self.backend = XPUAttentionBackend.__new__(XPUAttentionBackend) + self.backend.is_encoder_decoder = True + # Deliberately scattered (non-page-aligned) slot indices: a paged kernel + # would mis-read these; the page_size=1 gather must be alignment-agnostic. + perm = torch.randperm(self.TOTAL_SLOTS) + self.k_flat = torch.randn( + self.TOTAL_SLOTS, self.H, self.D, dtype=torch.bfloat16, device=self.dev + )[perm].contiguous() + self.v_flat = torch.randn( + self.TOTAL_SLOTS, self.H, self.D, dtype=torch.bfloat16, device=self.dev + )[perm].contiguous() + + def _check(self, *, cache_seqlens, cu_seqlens_q, causal): + cache_seqlens = cache_seqlens.to(self.dev) + cu_seqlens_q = cu_seqlens_q.to(self.dev) + num_rows = int(cu_seqlens_q[-1]) + m = int(cache_seqlens.max()) + # Rows packed valid-first; scatter distinct slots per request (build the + # permutation on CPU, then move -- randperm(device="xpu") is unreliable). + page_table = ( + torch.stack( + [ + torch.randperm(self.TOTAL_SLOTS)[:m] + for _ in range(cache_seqlens.numel()) + ] + ) + .to(torch.int32) + .to(self.dev) + ) + q = torch.randn(num_rows, self.H, self.D, dtype=torch.bfloat16, device=self.dev) + layer = SimpleNamespace( + is_cross_attention=not causal, + tp_q_head_num=self.H, + tp_k_head_num=self.H, + tp_v_head_num=self.H, + head_dim=self.D, + scaling=0.5, + logit_cap=0.0, + ) + key_cache = self.k_flat.view(-1, 1, self.H, self.D) + value_cache = self.v_flat.view(-1, 1, self.H, self.D) + + # causal=True mirrors decoder self-attn, causal=False cross-attn; the + # generic helper takes the (page_table, cache_seqlens, causal) that the + # caller's _encoder_decoder_page_table dispatch would have selected. + got = self.backend._forward_attn_flat_page_table( + q=q, + key_cache=key_cache, + value_cache=value_cache, + layer=layer, + page_table=page_table, + cache_seqlens=cache_seqlens, + cu_seqlens_q=cu_seqlens_q, + max_seqlen_q=1, + causal=causal, + ) + torch.xpu.synchronize() + want = _sdpa_ref( + q=q, + k_flat=self.k_flat, + v_flat=self.v_flat, + page_table=page_table, + cache_seqlens=cache_seqlens, + cu_seqlens_q=cu_seqlens_q, + scale=layer.scaling, + causal=causal, + ) + self.assertEqual(tuple(got.shape), (num_rows, self.H, self.D)) + self.assertTrue(torch.isfinite(got).all(), "attention output must be finite") + # bf16 kernel vs fp32 reference: loose tolerance. + torch.testing.assert_close(got.float(), want, rtol=2e-2, atol=2e-2) + + def test_cross_attention_decode_on_xpu(self): + # 1 query/request, attend all encoder KV, non-causal, unequal lengths. + self._check( + cache_seqlens=torch.tensor([5, 8], dtype=torch.int32), + cu_seqlens_q=torch.tensor([0, 1, 2], dtype=torch.int32), + causal=False, + ) + + def test_decoder_self_attention_decode_on_xpu(self): + self._check( + cache_seqlens=torch.tensor([4, 6], dtype=torch.int32), + cu_seqlens_q=torch.tensor([0, 1, 2], dtype=torch.int32), + causal=True, + ) + + def test_mixed_empty_batch_on_xpu(self): + # Mixed batch: request 0 has no keys (cache_seqlens==0), request 1 has some. + # On the real kernel the empty request's rows come back NaN/inf, so the + # backend must zero them without corrupting request 1. The SDPA oracle + # yields zeros for the empty request (empty-key contraction), so the shared + # assert_close plus the finiteness check guard against a regression that + # drops the zeroing and leaks NaN into the output. + self._check( + cache_seqlens=torch.tensor([0, 6], dtype=torch.int32), + cu_seqlens_q=torch.tensor([0, 1, 2], dtype=torch.int32), + causal=False, + ) + + +if __name__ == "__main__": + unittest.main()