flashinfer.kda_training.RecurrentKDATrainingContext

class flashinfer.kda_training.RecurrentKDATrainingContext(state_checkpoints: ~torch.Tensor, beta_active: ~torch.Tensor, _route: ~flashinfer._kda_training_impl._TrainingRouteSpec, _shape: ~flashinfer._kda_training_impl._TrainingShape, _q: ~torch.Tensor, _k: ~torch.Tensor, _v: ~torch.Tensor, _g: ~torch.Tensor, _beta: ~torch.Tensor, _A_log: ~torch.Tensor, _dt_bias: ~torch.Tensor, _initial_state: ~torch.Tensor, _cu_seqlens: ~torch.Tensor, _route_tensors: dict[str, ~torch.Tensor], _metadata: dict[str, object], _final_output_scratch: ~torch.Tensor, _final_descriptor_storage: ~torch.Tensor, _final_tensormap_workspace: ~torch.Tensor, _dummy_f32: ~torch.Tensor, _dummy_i32: ~torch.Tensor, _final_grid_ctas: int, _stream_ptr: int, _input_tensors: tuple[~torch.Tensor, ...] = (), _input_signatures: tuple[tuple, ...] = (), _saved_context_signatures: tuple[tuple, ...] = (), _final_descriptor_signature: tuple | None = None, _route_descriptor_signature: tuple | None = None, _backward_buffers: dict[str, ~torch.Tensor] = <factory>, _lock: ~_thread.allocate_lock = <factory>)

Route-tagged forward tapes consumed by the paired backward.

__init__(state_checkpoints: ~torch.Tensor, beta_active: ~torch.Tensor, _route: ~flashinfer._kda_training_impl._TrainingRouteSpec, _shape: ~flashinfer._kda_training_impl._TrainingShape, _q: ~torch.Tensor, _k: ~torch.Tensor, _v: ~torch.Tensor, _g: ~torch.Tensor, _beta: ~torch.Tensor, _A_log: ~torch.Tensor, _dt_bias: ~torch.Tensor, _initial_state: ~torch.Tensor, _cu_seqlens: ~torch.Tensor, _route_tensors: dict[str, ~torch.Tensor], _metadata: dict[str, object], _final_output_scratch: ~torch.Tensor, _final_descriptor_storage: ~torch.Tensor, _final_tensormap_workspace: ~torch.Tensor, _dummy_f32: ~torch.Tensor, _dummy_i32: ~torch.Tensor, _final_grid_ctas: int, _stream_ptr: int, _input_tensors: tuple[~torch.Tensor, ...] = (), _input_signatures: tuple[tuple, ...] = (), _saved_context_signatures: tuple[tuple, ...] = (), _final_descriptor_signature: tuple | None = None, _route_descriptor_signature: tuple | None = None, _backward_buffers: dict[str, ~torch.Tensor] = <factory>, _lock: ~_thread.allocate_lock = <factory>) None

Methods

__init__(state_checkpoints, beta_active, ...)

Attributes

state_checkpoints

beta_active