core.transformer.streamwise_residual_ops#
Full-width streamwise residual operations.
The pure PyTorch functions are the numerical reference and CPU fallback. CUDA training and inference can use dedicated Triton kernels that consume the padded controller logits directly. Those kernels form sigmoid factors in registers and compute controller gradients from per-program partials, avoiding both dense slot maps and a second activation-reading GEMM/BMM in backward.
BATCH is a runtime kernel argument rather than a tl.constexpr: it is only
compared against to mask the tail, so specializing on it buys nothing and would
force a separate JIT compile per token count. Dynamic-batching inference captures
one CUDA graph per token bucket, and a compile that misses graph warmup would
land inside stream capture.
Module Contents#
Classes#
Autograd wrapper for direct raw-logit streamwise Triton read. |
|
Autograd wrapper for direct raw-logit streamwise Triton writeback. |
Functions#
Flatten all leading dimensions for the two-dimensional Triton kernels. |
|
Restore leading dimensions after a Triton kernel launch. |
|
Use the block geometry tuned for the existing small-map MHC kernels. |
|
Return whether a direct raw-logit streamwise Triton kernel is supported. |
|
View the final dimension as |
|
Validate a vector containing one scalar controller per full-width stream. |
|
Cast and view stream factors for broadcasting over leading and hidden dimensions. |
|
Read |
|
Write one branch update independently into |
|
Reduce tile-local FP32 controller gradients without rereading activations. |
|
Return whether the fused streamwise operation needs backward state. |
|
Return whether an inference engine is driving a forward-only call. |
|
Read full-width streams from raw padded logits with fused CUDA dispatch. |
|
Write full-width streams from raw padded logits with fused CUDA dispatch. |
Data#
API#
- core.transformer.streamwise_residual_ops._STREAMWISE_MIN_BATCH#
256
- core.transformer.streamwise_residual_ops._STREAMWISE_MAX_STREAMS#
16
- core.transformer.streamwise_residual_ops._STREAMWISE_MAX_REDUCTION_BLOCK#
16384
- core.transformer.streamwise_residual_ops._flatten_leading(
- tensor: torch.Tensor,
Flatten all leading dimensions for the two-dimensional Triton kernels.
- core.transformer.streamwise_residual_ops._unflatten_leading(
- tensor: torch.Tensor,
- leading: torch.Size,
Restore leading dimensions after a Triton kernel launch.
- core.transformer.streamwise_residual_ops._streamwise_block_config(
- batch: int,
- stream_width: int,
Use the block geometry tuned for the existing small-map MHC kernels.
- core.transformer.streamwise_residual_ops._num_gradient_partials(batch: int, stream_width: int) int#
- core.transformer.streamwise_residual_ops._validate_logits(
- logits: torch.Tensor,
- num_streams: int,
- *,
- tensor_name: str,
- core.transformer.streamwise_residual_ops._validate_raw_read_inputs(
- hidden_states: torch.Tensor,
- read_logits: torch.Tensor,
- num_streams: int,
- core.transformer.streamwise_residual_ops._validate_raw_write_inputs(
- residual_stream: torch.Tensor,
- branch_update: torch.Tensor,
- write_logits: torch.Tensor,
- num_streams: int,
- retention_logits: torch.Tensor | None,
- core.transformer.streamwise_residual_ops._can_use_streamwise_triton(
- tensor: torch.Tensor,
- logits: torch.Tensor,
- num_streams: int,
- stream_width: int,
- *,
- enforce_training_limits: bool = True,
Return whether a direct raw-logit streamwise Triton kernel is supported.
enforce_training_limits=Falsedrops two training-path constraints. The minimum batch is a performance floor for amortizing fused forward and backward overhead, while the reduction-block cap bounds the controller-gradient partial reduction. A forward-only inference call allocates no gradient partials and runs no backward reduction, so neither limit needs to apply.Training keeps the limits even when grad is disabled. Residual-stream recompute replays the forward under
torch.no_grad()during backward, and that replay has to reproduce the original forward’s arithmetic exactly, so the kernel choice must not depend on whether grad happens to be enabled.
- core.transformer.streamwise_residual_ops._view_streams(
- tensor: torch.Tensor,
- num_streams: int,
- *,
- tensor_name: str,
View the final dimension as
num_streamscontiguous full-width streams.
- core.transformer.streamwise_residual_ops._validate_factors(factors: torch.Tensor, *, factor_name: str) int#
Validate a vector containing one scalar controller per full-width stream.
- core.transformer.streamwise_residual_ops._broadcast_factors(
- factors: torch.Tensor,
- streams: torch.Tensor,
Cast and view stream factors for broadcasting over leading and hidden dimensions.
- core.transformer.streamwise_residual_ops.streamwise_read(
- hidden_states: torch.Tensor,
- read_factors: torch.Tensor,
Read
Kfull-width streams into one branch-width activation.Given
hidden_states[..., k, :] = X_kand one scalarc_kper stream, this computessum_k c_k X_k. Native autograd supplies activation and factor gradients without constructing a masked slot-mixing matrix.
- core.transformer.streamwise_residual_ops.streamwise_writeback(
- residual_stream: torch.Tensor,
- branch_update: torch.Tensor,
- write_factors: torch.Tensor,
- *,
- retention_factors: torch.Tensor | None = None,
Write one branch update independently into
Kfull-width streams.For stream
k, this computesY_k = gamma_k X_k + w_k U. Omittingretention_factorsgives identity carry,gamma_k = 1. The returned tensor owns the only full-width output allocation; no expanded update tensor or masked slot map is materialized.
- core.transformer.streamwise_residual_ops._reduce_streamwise_partials(
- partials: torch.Tensor,
- grad_logits: torch.Tensor,
- *,
- second_partials: torch.Tensor | None = None,
- second_grad_logits: torch.Tensor | None = None,
Reduce tile-local FP32 controller gradients without rereading activations.
- core.transformer.streamwise_residual_ops._streamwise_sigmoid_read_triton(
- hidden_states: torch.Tensor,
- read_logits: torch.Tensor,
- num_streams: int,
- core.transformer.streamwise_residual_ops._streamwise_sigmoid_read_backward_triton(
- hidden_states: torch.Tensor,
- grad_output: torch.Tensor,
- read_logits: torch.Tensor,
- num_streams: int,
- core.transformer.streamwise_residual_ops._streamwise_sigmoid_write_triton(
- residual_stream: torch.Tensor,
- branch_update: torch.Tensor,
- write_logits: torch.Tensor,
- num_streams: int,
- retention_logits: torch.Tensor | None,
- retention_max_forget: float,
- core.transformer.streamwise_residual_ops._streamwise_sigmoid_write_backward_triton(
- residual_stream: torch.Tensor | None,
- branch_update: torch.Tensor,
- grad_output: torch.Tensor,
- write_logits: torch.Tensor,
- num_streams: int,
- retention_logits: torch.Tensor | None,
- retention_max_forget: float,
- class core.transformer.streamwise_residual_ops._StreamwiseSigmoidRead#
Bases:
torch.autograd.FunctionAutograd wrapper for direct raw-logit streamwise Triton read.
- static forward(
- ctx,
- hidden_states: torch.Tensor,
- read_logits: torch.Tensor,
- num_streams: int,
Run the fused streamwise read and save tensors for backward.
- static backward(
- ctx,
- grad_output: torch.Tensor,
Compute activation and read-controller gradients.
- class core.transformer.streamwise_residual_ops._StreamwiseSigmoidWriteback#
Bases:
torch.autograd.FunctionAutograd wrapper for direct raw-logit streamwise Triton writeback.
- static forward(
- ctx,
- residual_stream: torch.Tensor,
- branch_update: torch.Tensor,
- write_logits: torch.Tensor,
- retention_logits: torch.Tensor | None,
- num_streams: int,
- retention_max_forget: float,
Run fused streamwise writeback and save tensors for backward.
- static backward(
- ctx,
- grad_output: torch.Tensor,
Compute residual, update, map, and retention gradients.
- core.transformer.streamwise_residual_ops._streamwise_autograd_needed(*tensors: torch.Tensor | None) bool#
Return whether the fused streamwise operation needs backward state.
- core.transformer.streamwise_residual_ops._streamwise_inference_forward(*tensors: torch.Tensor | None) bool#
Return whether an inference engine is driving a forward-only call.
Gating on the engine flag rather than on grad state alone keeps training self-consistent: residual-stream recompute replays the forward under
torch.no_grad(), and it must select the same kernel as the original grad-enabled forward.
- core.transformer.streamwise_residual_ops.streamwise_sigmoid_read(
- hidden_states: torch.Tensor,
- read_logits: torch.Tensor,
- num_streams: int,
Read full-width streams from raw padded logits with fused CUDA dispatch.
- core.transformer.streamwise_residual_ops.streamwise_sigmoid_writeback(
- residual_stream: torch.Tensor,
- branch_update: torch.Tensor,
- write_logits: torch.Tensor,
- num_streams: int,
- *,
- retention_logits: torch.Tensor | None = None,
- retention_max_forget: float = 0.0,
Write full-width streams from raw padded logits with fused CUDA dispatch.