跳到正文
HNHacker News·

PSSA: A non-transformer language model written from scratch in Rust

AI 摘要

PSSA, a plastic state-space architecture, is a non-transformer language model written in Rust. It demonstrates superior performance compared to transformers on a held-out slice of 198,939 unseen tokens, achieving a cross-entropy of 3.997 versus 4.429, a perplexity of 54.4 versus 83.8, and a next-token accuracy of 24.1% versus 18.0%. This indicates better generalization rather than harder memorization. Contributions are welcome, particularly for kernel performance and recurrent baselines.

时间与来源

时间显示为 UTC

显示时区:UTC

本地时区尚不可用,暂时显示 UTC。

发布当时偏移:UTC+02026年9月30日 03:19 UTC

收录当时偏移:UTC+02026年9月30日 04:00 UTC

发布
2026年9月30日 03:19
收录
2026年9月30日 04:00
来源类型
开发者社区
档位
社区
信源状态
正常

档位是按信源手工设定的编辑判断,不是逐条打分。

PSSA: a plastic state-space architecture

PSSA is a small language model that is not a transformer. It reads text one token at a time through a recurrent state-space layer, keeps a bank of episodic memories it can look things up in, and rewrites part of its own weights while it runs. It is written in Rust from scratch, with no PyTorch, no TensorFlow, and no ML framework of any kind underneath it.

At matched parameters and on the same corpus, it learns faster than a transformer and generates text about twelve times quicker on the same CPU.

Why Rust, and why that is not the point

Not for speed points, and not because the language makes the architecture better. PSSA needed per-token weight updates, a memory bank written during the forward pass, and a scalar reference path that every batched kernel could be differentiated against. Expressing that inside an autograd framework meant fighting the framework at every step, so the linear algebra is written directly instead. That made the plastic parts straightforward and the gradients checkable against a reference to around 3e-8. The architecture is the claim here. The implementation language is a detail, and a Python port is welcome.

How it differs from a transformer

A transformer scores every pair of tokens in the context, so its cost per step grows with the square of the sequence length and the whole context is re-read at every step. PSSA carries one fixed-size state along the sequence in a single left-to-right pass, and looks things up in a memory bank instead of re-reading the context, so cost grows linearly with length.

The model

Every token goes through one PSSA layer: a selective state-space recurrence, a bounded read from an episodic memory bank in hyperbolic space, a learned gate that decides how much of that read reaches the residual stream, and a SiLU MLP. The defaults are d_m = 256 channels, d_s = 16 states per channel, and a rank-16 adapter.

The recurrence

Write x for the layer-normalized token embedding. Three projections are read off the token itself, which is what makes the recurrence selective rather than fixed:

delta = softplus(W_delta x) per-channel step size, delta in R^d_m B = W_B x input map, B in R^d_s C = W_C x output map, C in R^d_s

The transition is diagonal, one rate per (channel, state) pair, kept negative by construction so the recurrence cannot blow up:

A = -softplus(A_raw) A in R^(d_m x d_s)

Discretizing that continuous system with step delta gives the per-token update. h carries across tokens and across chunk boundaries during training:

Abar_ij = exp(delta_i * A_ij) Bbar_ij = delta_i * B_j

h_ij

A_raw is initialized so each channel's 16 rates sit on log-spaced timescales tau from 1.5 to 200 tokens, in the spirit of the HiPPO initialization. A single channel therefore starts out holding the last two tokens and the last two hundred at the same time, and training moves those horizons rather than discovering them from scratch.

This half of the layer is a selective diagonal SSM and claims no novelty; it is the same family as S4 and Mamba, written out scalar-first so the backward pass can be checked term by term.

The memory read

The part that is specific to PSSA is what happens to y. A query is formed from both the current token and the current state, so retrieval is conditioned on where the recurrence has got to and not only on the token in hand:

q = W_qx x + W_qh y qh = proj(q) diffeomorphic map into the Poincare ball, |qh|

The read is bounded at four slots, weighted by a softmax over hyperbolic distance at temperature tau_mem:

w = softmax(-d_H(qh, k_s) / tau_mem) over the 4 nearest slots m = sum_k w_k * v_k

来源·Hacker News·github.com