Skip to content

The examples mentioned on https://pypi.org/project/xlstm do not work #107

Description

@roman8ivanov

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

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions