Hi.
I use CUDA 12.8, pytorch 2.8.0, triton 3.4 on Windows. I tried:
import torch
from xlstm.xlstm_large.model import xLSTMLargeConfig, xLSTMLarge
configure the model with TFLA Triton kernels
xlstm_config = xLSTMLargeConfig(
embedding_dim=512,
num_heads=4,
num_blocks=6,
vocab_size=2048,
return_last_states=True,
mode="inference",
chunkwise_kernel="chunkwise--triton_xl_chunk", # xl_chunk == TFLA kernels
sequence_kernel="native_sequence__triton",
step_kernel="triton",
)
instantiate the model
xlstm = xLSTMLarge(xlstm_config)
xlstm = xlstm.to("cuda")
create inputs
input = torch.randint(0, 2048, (3, 256)).to("cuda")
run a forward pass
out = xlstm(input)
I have got the below mentioned error executing out = xlstm(input):
CompilationError: at 10:23:
def flip(x, dim=None):
"""
Flips a tensor x along the dimension dim.
:param x: the first input tensor
:type x: Block
:param dim: the dimension to flip along
:type dim: int
"""
core.static_assert(-len(x.shape) <= dim and dim < len(x.shape))
^
TypeError("'<=' not supported between instances of 'int' and 'NoneType'")
The above exception was the direct cause of the following exception:
Traceback (most recent call last):
Cell In[10], line 1
out = xlstm(input)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1773 in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1784 in _call_impl
return forward_call(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\xlstm\xlstm_large\model.py:146 in forward
x, state = self.backbone(x, state)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1773 in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1784 in _call_impl
return forward_call(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\xlstm\xlstm_large\model.py:217 in forward
x, block_state_new = block(x, block_state)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1773 in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1784 in _call_impl
return forward_call(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\xlstm\xlstm_large\model.py:503 in forward
x_mlstm, state = self.mlstm_layer(x_mlstm, state)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1773 in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1784 in _call_impl
return forward_call(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\xlstm\xlstm_large\model.py:425 in forward
h, state = self.mlstm_backend(
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1773 in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1784 in _call_impl
return forward_call(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\backend_module.py:207 in forward
return self._inference_fn(
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\kernel_wrappers.py:131 in wrap_chunkwise__arbitrary_sequence_length
h_out, (c_state, n_state, m_state) = mlstm_chunkwise_kernel(
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\chunkwise\triton_xl_chunk\fwbw.py:272 in mlstm_chunkwise__xl_chunk
matH_out, matC_last, vecN_last, scaM_last = _mlstm_chunkwise_fwbw.apply(
File C:\ProgramData\anaconda3\Lib\site-packages\torch\autograd\function.py:576 in apply
return super().apply(*args, **kwargs) # type: ignore[misc]
File C:\ProgramData\anaconda3\Lib\site-packages\torch\amp\autocast_mode.py:528 in decorate_fwd
return fwd(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\utils.py:33 in wrapper
return fn(
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\chunkwise\triton_xl_chunk\fwbw.py:51 in forward
matH_out, vecN_out, vecM_out, last_states, all_states = mlstm_chunkwise_fw(
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\utils.py:48 in wrapper
return fn(
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\chunkwise\triton_xl_chunk\fw.py:74 in mlstm_chunkwise_fw
matC_k_states, vecN_k_states, scaMinter_k_states = mlstm_chunkwise__recurrent_fw_C(
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\chunkwise\triton_xl_chunk\fw_recurrent.py:97 in mlstm_chunkwise__recurrent_fw_C
mlstm_chunkwise__recurrent_fw_C_kernel[grid](
File C:\ProgramData\anaconda3\Lib\site-packages\triton\runtime\jit.py:390 in
return lambda *args, **kwargs: self.run(grid=grid, warmup=False, *args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\triton\runtime\jit.py:594 in run
kernel = self.compile(src, target=target, options=options.dict)
File C:\ProgramData\anaconda3\Lib\site-packages\triton\compiler\compiler.py:339 in compile
module = src.make_ir(options, codegen_fns, module_map, context)
File C:\ProgramData\anaconda3\Lib\site-packages\triton\compiler\compiler.py:83 in make_ir
return ast_to_ttir(self.fn, self, context=context, options=options, codegen_fns=codegen_fns,
CompilationError: at 153:39:
).to(tl.float32)
vecFlogsig_k_val = tl.log(tl.sigmoid(vecF_k_val))
vecFlogsig_masked = tl.where(idx_L < L - 1, vecFlogsig_k_val, 0.0).to(
tl.float32
)
vecI_k_val = tl.load(vecI + idx_b_BNH * str_vecFI_B_NH + k * L + idx_L).to(
tl.float32
)
vecA_k_val = tl.flip(tl.cumsum(tl.flip(vecFlogsig_masked), axis=0)) + vecI_k_val
Hi.
I use CUDA 12.8, pytorch 2.8.0, triton 3.4 on Windows. I tried:
import torch
from xlstm.xlstm_large.model import xLSTMLargeConfig, xLSTMLarge
configure the model with TFLA Triton kernels
xlstm_config = xLSTMLargeConfig(
embedding_dim=512,
num_heads=4,
num_blocks=6,
vocab_size=2048,
return_last_states=True,
mode="inference",
chunkwise_kernel="chunkwise--triton_xl_chunk", # xl_chunk == TFLA kernels
sequence_kernel="native_sequence__triton",
step_kernel="triton",
)
instantiate the model
xlstm = xLSTMLarge(xlstm_config)
xlstm = xlstm.to("cuda")
create inputs
input = torch.randint(0, 2048, (3, 256)).to("cuda")
run a forward pass
out = xlstm(input)
I have got the below mentioned error executing out = xlstm(input):
CompilationError: at 10:23:
def flip(x, dim=None):
"""
Flips a tensor
xalong the dimensiondim.TypeError("'<=' not supported between instances of 'int' and 'NoneType'")
The above exception was the direct cause of the following exception:
Traceback (most recent call last):
Cell In[10], line 1
out = xlstm(input)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1773 in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1784 in _call_impl
return forward_call(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\xlstm\xlstm_large\model.py:146 in forward
x, state = self.backbone(x, state)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1773 in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1784 in _call_impl
return forward_call(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\xlstm\xlstm_large\model.py:217 in forward
x, block_state_new = block(x, block_state)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1773 in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1784 in _call_impl
return forward_call(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\xlstm\xlstm_large\model.py:503 in forward
x_mlstm, state = self.mlstm_layer(x_mlstm, state)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1773 in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1784 in _call_impl
return forward_call(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\xlstm\xlstm_large\model.py:425 in forward
h, state = self.mlstm_backend(
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1773 in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\torch\nn\modules\module.py:1784 in _call_impl
return forward_call(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\backend_module.py:207 in forward
return self._inference_fn(
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\kernel_wrappers.py:131 in wrap_chunkwise__arbitrary_sequence_length
h_out, (c_state, n_state, m_state) = mlstm_chunkwise_kernel(
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\chunkwise\triton_xl_chunk\fwbw.py:272 in mlstm_chunkwise__xl_chunk
matH_out, matC_last, vecN_last, scaM_last = _mlstm_chunkwise_fwbw.apply(
File C:\ProgramData\anaconda3\Lib\site-packages\torch\autograd\function.py:576 in apply
return super().apply(*args, **kwargs) # type: ignore[misc]
File C:\ProgramData\anaconda3\Lib\site-packages\torch\amp\autocast_mode.py:528 in decorate_fwd
return fwd(*args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\utils.py:33 in wrapper
return fn(
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\chunkwise\triton_xl_chunk\fwbw.py:51 in forward
matH_out, vecN_out, vecM_out, last_states, all_states = mlstm_chunkwise_fw(
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\utils.py:48 in wrapper
return fn(
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\chunkwise\triton_xl_chunk\fw.py:74 in mlstm_chunkwise_fw
matC_k_states, vecN_k_states, scaMinter_k_states = mlstm_chunkwise__recurrent_fw_C(
File C:\ProgramData\anaconda3\Lib\site-packages\mlstm_kernels\torch\chunkwise\triton_xl_chunk\fw_recurrent.py:97 in mlstm_chunkwise__recurrent_fw_C
mlstm_chunkwise__recurrent_fw_C_kernel[grid](
File C:\ProgramData\anaconda3\Lib\site-packages\triton\runtime\jit.py:390 in
return lambda *args, **kwargs: self.run(grid=grid, warmup=False, *args, **kwargs)
File C:\ProgramData\anaconda3\Lib\site-packages\triton\runtime\jit.py:594 in run
kernel = self.compile(src, target=target, options=options.dict)
File C:\ProgramData\anaconda3\Lib\site-packages\triton\compiler\compiler.py:339 in compile
module = src.make_ir(options, codegen_fns, module_map, context)
File C:\ProgramData\anaconda3\Lib\site-packages\triton\compiler\compiler.py:83 in make_ir
return ast_to_ttir(self.fn, self, context=context, options=options, codegen_fns=codegen_fns,
CompilationError: at 153:39:
).to(tl.float32)