↓ 1 callersFunctioncount_flops_mlstm_v1_block_fw(
S,
d,
dqk,
dv,
Nh,
chunk_size=64,
pf_ffn=4,
factor_sig=1,
factor_exp=1,
agents/backbones/xlstm/mlstm_kernels/utils/flops/mlstm_block_flop_counts.py:94
↓ 1 callersFunctioncount_flops_mlstm_v1_layer_fw(
S,
d,
dqk,
dv,
Nh,
chunk_size,
factor_sig=1,
factor_exp=1,
factor_max=1,
agents/backbones/xlstm/mlstm_kernels/utils/flops/mlstm_block_flop_counts.py:53
↓ 1 callersFunctioncount_flops_mlstm_v2_block_fw(
S,
d,
dqk,
dv,
Nh,
qk_block_size=4,
qk_pf=1,
v_block_size=4,
v_pf=1,
agents/backbones/xlstm/mlstm_kernels/utils/flops/mlstm_block_flop_counts.py:135
↓ 1 callersFunctioncount_flops_slstm_block_fw(
S, d, Nh, conv1d_kernel_size=4, pf_ffn=1.3, factor_exp=1, count_ln_flops: Callable[[int], int] = _count_
agents/backbones/xlstm/mlstm_kernels/utils/flops/slstm_block_flop_counts.py:26
↓ 1 callersFunctioncreate_block(
d_model,
d_intermediate,
ssm_cfg=None,
attn_layer_idx=None,
attn_cfg=None,
norm_epsi
agents/backbones/mamba/mamba.py:30
↓ 1 callersMethoddpm_solver_adaptive(self, state, action, goal, t_start, t_end, order=3, rtol=0.05, atol=0.0078, h_init=0.05, pcoeff=0., icoeff=1
agents/models/beso/models/edm_diffusion/gc_sampling.py:618
↓ 1 callersMethoddpm_solver_fast(self, state, action, goal, t_start, t_end, nfe, eta=0., s_noise=1., noise_sampler=None)
agents/models/beso/models/edm_diffusion/gc_sampling.py:579
↓ 1 callersMethodlabel_seq(self, collected_data, dataset, seq_length, idx, task)
agents/models/beso/utils/automatic_lang_annotator_mp.py:311
↓ 1 callersFunctionmlstm_chunkwise__parallel_bw_dQKV(
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/bw_parallel.py:12
↓ 1 callersFunctionmlstm_chunkwise__parallel_bw_dQKV Computes the gradients for the query, key and value matrices. Args: matQ: Tensor containing the query vectors. Shape (B, NH, S, DHQK
agents/backbones/xlstm/mlstm_kernels/jax/chunkwise/triton_limit_chunk/bw_parallel.py:17
↓ 1 callersFunctionmlstm_chunkwise__parallel_fw_H(
matQ: torch.Tensor, # (B, NH, S, DHQK)
matK: torch.Tensor, # (B, NH, S, DHQK)
matV: torch.Tens
agents/backbones/xlstm/mlstm_kernels/torch/chunkwise/native/fw.py:131