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#

_StreamwiseSigmoidRead

Autograd wrapper for direct raw-logit streamwise Triton read.

_StreamwiseSigmoidWriteback

Autograd wrapper for direct raw-logit streamwise Triton writeback.

Functions#

_flatten_leading

Flatten all leading dimensions for the two-dimensional Triton kernels.

_unflatten_leading

Restore leading dimensions after a Triton kernel launch.

_streamwise_block_config

Use the block geometry tuned for the existing small-map MHC kernels.

_num_gradient_partials

_validate_logits

_validate_raw_read_inputs

_validate_raw_write_inputs

_can_use_streamwise_triton

Return whether a direct raw-logit streamwise Triton kernel is supported.

_view_streams

View the final dimension as num_streams contiguous full-width streams.

_validate_factors

Validate a vector containing one scalar controller per full-width stream.

_broadcast_factors

Cast and view stream factors for broadcasting over leading and hidden dimensions.

streamwise_read

Read K full-width streams into one branch-width activation.

streamwise_writeback

Write one branch update independently into K full-width streams.

_reduce_streamwise_partials

Reduce tile-local FP32 controller gradients without rereading activations.

_streamwise_sigmoid_read_triton

_streamwise_sigmoid_read_backward_triton

_streamwise_sigmoid_write_triton

_streamwise_sigmoid_write_backward_triton

_streamwise_autograd_needed

Return whether the fused streamwise operation needs backward state.

_streamwise_inference_forward

Return whether an inference engine is driving a forward-only call.

streamwise_sigmoid_read

Read full-width streams from raw padded logits with fused CUDA dispatch.

streamwise_sigmoid_writeback

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,
) → tuple[torch.Tensor, torch.Size]#

Flatten all leading dimensions for the two-dimensional Triton kernels.

core.transformer.streamwise_residual_ops._unflatten_leading(
tensor: torch.Tensor,
leading: torch.Size,
) → torch.Tensor#

Restore leading dimensions after a Triton kernel launch.

core.transformer.streamwise_residual_ops._streamwise_block_config(
batch: int,
stream_width: int,
) → tuple[int, int, 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,
) → None#
core.transformer.streamwise_residual_ops._validate_raw_read_inputs(
hidden_states: torch.Tensor,
read_logits: torch.Tensor,
num_streams: int,
) → 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,
) → int#
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,
) → bool#

Return whether a direct raw-logit streamwise Triton kernel is supported.

enforce_training_limits=False drops 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,
) → torch.Tensor#

View the final dimension as num_streams contiguous 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,
) → 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,
) → torch.Tensor#

Read K full-width streams into one branch-width activation.

Given hidden_states[..., k, :] = X_k and one scalar c_k per stream, this computes sum_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,
) → torch.Tensor#

Write one branch update independently into K full-width streams.

For stream k, this computes Y_k = gamma_k X_k + w_k U. Omitting retention_factors gives 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,
) → 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,
) → torch.Tensor#
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,
) → tuple[torch.Tensor, torch.Tensor]#
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,
) → torch.Tensor#
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,
) → tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]#
class core.transformer.streamwise_residual_ops._StreamwiseSigmoidRead#

Bases: torch.autograd.Function

Autograd wrapper for direct raw-logit streamwise Triton read.

static forward(
ctx,
hidden_states: torch.Tensor,
read_logits: torch.Tensor,
num_streams: int,
) → torch.Tensor#

Run the fused streamwise read and save tensors for backward.

static backward(
ctx,
grad_output: torch.Tensor,
) → tuple[torch.Tensor | None, ...]#

Compute activation and read-controller gradients.

class core.transformer.streamwise_residual_ops._StreamwiseSigmoidWriteback#

Bases: torch.autograd.Function

Autograd 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,
) → torch.Tensor#

Run fused streamwise writeback and save tensors for backward.

static backward(
ctx,
grad_output: torch.Tensor,
) → tuple[torch.Tensor | None, ...]#

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,
) → torch.Tensor#

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,
) → torch.Tensor#

Write full-width streams from raw padded logits with fused CUDA dispatch.