↓ 2 callersMethodannotate(self, episode, dataset, collected_data, global_task_counter, num_samples)
agents/models/beso/utils/automatic_lang_annotator_mp.py:262
↓ 2 callersMethodcheck_done(self, counter, num_samples, batch_idx, num_batches, mode)
agents/models/beso/utils/automatic_lang_annotator_mp.py:237
↓ 2 callersFunctioncompute_chunkwise_log_gates_vecB_vecA(
vecI: jax.Array, # (B, NH, S)
vecF: jax.Array, # (B, NH, S)
chunk_size: int,
return_vecB_o
agents/backbones/xlstm/mlstm_kernels/jax/chunkwise/triton_xl_chunk/chunkwise_gates.py:13
↓ 2 callersMethoddpm_solver_1_step(self, state, action, goal, t, t_next, eps_cache=None)
agents/models/beso/models/edm_diffusion/gc_sampling.py:549
↓ 2 callersMethoddpm_solver_3_step(self, state, action, goal, t, t_next, r1=1 / 3, r2=2 / 3, eps_cache=None)
agents/models/beso/models/edm_diffusion/gc_sampling.py:566
↓ 2 callersMethodforward_enc_only(self, states, actions=None, goals=None, sigma=None, uncond: Optional[bool] = False)
agents/models/beso/models/networks/mdtv_transformer.py:213
↓ 2 callersFunctionmlstm_chunkwise__recurrent_fw_C(
matK: torch.Tensor, # (B, NH, S, DHQK)
matV: torch.Tensor, # (B, NH, S, DHHV)
vecB: torch.Tens
agents/backbones/xlstm/mlstm_kernels/torch/chunkwise/triton_limit_chunk/fw_recurrent.py:12
↓ 2 callersFunctionmlstm_chunkwise__recurrent_fw_C(
matK: torch.Tensor, # (B, NH, S, DHQK)
matV: torch.Tensor, # (B, NH, S, DHHV)
vecF: torch.Tens
agents/backbones/xlstm/mlstm_kernels/torch/chunkwise/triton_xl_chunk/fw_recurrent.py:12
↓ 2 callersFunctionmlstm_chunkwise__recurrent_fw_C(
matK: torch.Tensor, # (B, NH, S, DHQK)
matV: torch.Tensor, # (B, NH, S, DHHV)
vecB: torch.Tens
agents/backbones/xlstm/mlstm_kernels/torch/chunkwise/native/fw.py:29
↓ 2 callersFunctionmlstm_parallel_fw(
matQ: jax.Array,
matK: jax.Array,
matV: jax.Array,
vecI: jax.Array,
vecF: jax.Array,
agents/backbones/xlstm/mlstm_kernels/jax/parallel/native_stablef/fw.py:15
↓ 2 callersFunctionmlstm_parallel_fw(
matQ: jax.Array,
matK: jax.Array,
matV: jax.Array,
vecI: jax.Array,
vecF: jax.Array,
agents/backbones/xlstm/mlstm_kernels/jax/parallel/native/fw.py:15
↓ 2 callersFunctionplot_numerical_diffs_per_batchhead(
baseline,
target=None,
title="",
vmin=0.0,
vmax=1e-2,
figsize=(10, 6),
rtol: flo
agents/backbones/xlstm/mlstm_kernels/utils/plot/diff_imshow.py:118