↓ 8 callersMethodeps(self, eps_cache, key, state, action, goal, t, *args, **kwargs)
agents/models/beso/models/edm_diffusion/gc_sampling.py:540
↓ 5 callersFunctioncount_flops_mlstm_chunkwise_fw(L, Nc, dqk, dv, Nh, factor_exp, factor_max, factor_mask)
agents/backbones/xlstm/mlstm_kernels/utils/flops/mlstm_block_flop_counts.py:34
↓ 5 callersMethodsample(self, z, state, latent_goal, null_cond=None, sample_steps=50, cfg=2.0)
agents/models/flow_matching/rf.py:40
↓ 3 callersFunction_mlstm_recurrent_sequence_loop_fw(
mlstm_step_fn: Callable,
matQ: torch.Tensor, # (B, NH, S, DHQK)
matK: torch.Tensor, # (B, NH,
agents/backbones/xlstm/mlstm_kernels/torch/recurrent/native_sequence.py:13
↓ 3 callersMethoddpm_solver_2_step(self, state, action, goal, t, t_next, r1=1 / 2, eps_cache=None)
agents/models/beso/models/edm_diffusion/gc_sampling.py:556
↓ 3 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/native/fw.py:224
↓ 2 callersMethod__init__(
self,
dim,
custom_freqs = None,
freqs_for = 'lang',
theta = 10000,
agents/models/beso/models/networks/transformers/position_embeddings.py:84
↓ 2 callersFunction_attn_bwd_dkdv(
dk,
dv, #
Q,
k,
v,
sm_scale, #
DO, #
M,
D, #
# shared by Q/K/V/D
agents/backbones/xlstm/mlstm_kernels/baselines/flash_attention/triton_tutorial.py:279
↓ 2 callersFunction_attn_bwd_dq(
dq,
q,
K,
V, #
do,
m,
D,
# shared by Q/K/V/DO.
stride_tok,
stride_d
agents/backbones/xlstm/mlstm_kernels/baselines/flash_attention/triton_tutorial.py:344
↓ 2 callersFunction_attn_fwd_inner(
acc,
l_i,
m_i,
q, #
K_block_ptr,
V_block_ptr, #
start_m,
qk_scale, #
agents/backbones/xlstm/mlstm_kernels/baselines/flash_attention/triton_tutorial.py:33