↓ 1 callersFunctionmlstm_chunkwise_bw(
## Forward arguments
matQ: torch.Tensor, # (B, NH, S, DHQK)
matK: torch.Tensor, # (B, NH, S, D
agents/backbones/xlstm/mlstm_kernels/torch/chunkwise/triton_limit_chunk/bw.py:15
↓ 1 callersFunctionmlstm_chunkwise_bw(
## Forward arguments
matQ: torch.Tensor, # (B, NH, S, DHQK)
matK: torch.Tensor, # (B, NH, S, D
agents/backbones/xlstm/mlstm_kernels/torch/chunkwise/triton_xl_chunk/bw.py:27
↓ 1 callersFunctionmlstm_chunkwise_bw(
## Forward arguments
matQ: torch.Tensor, # (B, NH, S, DHQK)
matK: torch.Tensor, # (B, NH, S, D
agents/backbones/xlstm/mlstm_kernels/torch/chunkwise/native/bw.py:194
↓ 1 callersFunctionmlstm_chunkwise_bw(
# Forward arguments
matQ: jax.Array, # (B, NH, S, DHQK)
matK: jax.Array, # (B, NH, S, DHQK)
agents/backbones/xlstm/mlstm_kernels/jax/chunkwise/triton_xl_chunk/bw.py:20
↓ 1 callersFunctionmlstm_chunkwise_fw(
matQ: torch.Tensor, # (B, NH, S, DHQK)
matK: torch.Tensor, # (B, NH, S, DHQK)
matV: torch.Tens
agents/backbones/xlstm/mlstm_kernels/torch/chunkwise/triton_limit_chunk/fw.py:13
↓ 1 callersFunctionmlstm_chunkwise_fw(
matQ: torch.Tensor, # (B, NH, S, DHQK)
matK: torch.Tensor, # (B, NH, S, DHQK)
matV: torch.Tens
agents/backbones/xlstm/mlstm_kernels/torch/chunkwise/triton_xl_chunk/fw.py:20
↓ 1 callersFunctionmlstm_recurrent_step__native_fwThis is a single step of the mLSTM operation in recurrent form. Args: matC_old: (B, NH, DHQK, DHV) vecN_old: (B, NH, DHQK)
agents/backbones/xlstm/mlstm_kernels/torch/recurrent/native_step.py:8
↓ 1 callersFunctionmlstm_recurrent_step__triton_alternate_fw(
matC_old: torch.Tensor, # (B, NH, DHQK, DHV)
vecN_old: torch.Tensor, # (B, NH, DHQK)
scaM_old:
agents/backbones/xlstm/mlstm_kernels/torch/recurrent/triton_step_alternate.py:18
↓ 1 callersFunctionmlstm_recurrent_step__triton_fw(
matC_old: torch.Tensor, # (B, NH, DHQK, DHHV)
vecN_old: torch.Tensor, # (B, NH, DHQK)
scaM_old
agents/backbones/xlstm/mlstm_kernels/torch/recurrent/triton_step.py:14
↓ 1 callersFunctionmlstm_recurrent_step__triton_fw(
matC_state: jax.Array, # (B, NH, DHQK, DHV)
vecN_state: jax.Array, # (B, NH, DHQK)
scaM_state:
agents/backbones/xlstm/mlstm_kernels/jax/recurrent/triton_step.py:15
↓ 1 callersFunctionnaive_recurrent_gla(q, k, v, gk, initial_state=None, output_final_state=False, causal=True)
agents/backbones/xlstm/mlstm_kernels/baselines/flash_linear_attention/gla/naive.py:13
↓ 1 callersFunctionplot_error_statistics_over_time_single(
errors: np.ndarray, # shape: (num_timesteps, num_features)
percentiles: list = [50, 90, 100],
t
agents/backbones/xlstm/mlstm_kernels/utils/plot/diff_lineplot.py:29
↓ 1 callersFunctionplot_numerical_diffs_single(
baseline,
target=None,
title="",
vmin=0.0,
vmax=1e-2,
figsize=(10, 6),
convert_t
agents/backbones/xlstm/mlstm_kernels/utils/plot/diff_imshow.py:98