↓ 1 callersMethod__init__(
self,
obs_dim: int,
goal_dim: int,
device: str,
n_obs_token: int,
agents/models/beso/models/networks/mdtv_transformer.py:38
↓ 1 callersMethod__init__(
self,
dim: int,
depth: int,
dim_head: int = 64,
heads: int = 8,
agents/models/beso/models/networks/transformers/perceiver_resampler.py:83
↓ 1 callersFunction_mlstm_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/native/bw.py:98
↓ 1 callersFunctioncompute_chunkwise_log_gates_vecB_vecA(
vecI: torch.Tensor, # (B, NH, S)
vecF: torch.Tensor, # (B, NH, S)
chunk_size: int,
)
agents/backbones/xlstm/mlstm_kernels/torch/chunkwise/triton_xl_chunk/chunkwise_gates.py:16
↓ 1 callersFunctioncompute_errors_per_batchhead(
baseline: np.ndarray, # (B, NH, S, ...)
target: np.ndarray, # (B, NH, S, ...)
)
agents/backbones/xlstm/mlstm_kernels/utils/plot/diff_lineplot.py:10
↓ 1 callersFunctioncount_flops_ffn_layer_fw(S, d, pf, factor_gelu=1, count_ln_flops: Callable[[int], int] = _count_ln_flops)
agents/backbones/xlstm/mlstm_kernels/utils/flops/mlstm_block_flop_counts.py:86