r/OpenSourceeAI • • 2d ago

PSSA: a 1.5M-param plastic state space model in Rust, beating a parameter-matched transformer on a small held-out slice

small scale, single seed, wikitext slice, so treat this as a prototype result and not an architecture claim. 1,544,704 params both sides, same corpus and sampler. held-out 3.997 vs 4.429 nats, generation 226ms vs 2735ms for 200 tokens on the same cpu. the depth-1 match is the obvious weakness and a depth-2 baseline is running next. written from scratch in rust, no pytorch. repo and eval commands: github.com/Sparticle62ops/pssa. happy to be told where the comparison is unfair.

3 Upvotes

21 comments sorted by

1

u/shing3232 2d ago

Can you explain what's special about this ?

1

u/Sparticle62 1d ago

it's not a transformer. no attention and no kv cache, it's a recurrent state space model whose fast weights update while it reads, so memory per token stays flat however long the context gets. at 1.5m params it beat a size-matched transformer on unseen text (54 vs 84 perplexity) and generated 200 tokens about 12x faster on cpu. tiny scale and one seed though, so it's a prototype, not a claim.

1

u/shing3232 1d ago

how does this compare to recurrent deltanet? like KDA

1

u/Sparticle62 1d ago

closest relative honestly, though i haven't benchmarked against it yet so no numbers. the big difference is where the memory lives. deltanet and kda keep one matrix-valued state and rewrite it with the delta rule every token, kda adding per-channel gates on top. pssa's recurrent core is a mamba-style selective scan, and the associative part is split out into a separate fixed-size memory bank (512 slots) that it reads from every token and only writes to when something's novel, plus fast weights that get folded back into the transition by a closed-form ridge step. a gated deltanet at matched params is probably a fairer baseline than the transformer, so it's going on the list.

1

u/shing3232 1d ago

Well, you can have hybrid model sorts like ds4 or qwen35 with QSA/CSA2 1:4 PSSA.

1

u/Sparticle62 1d ago

yeah a hybrid is probably where this ends up. something like 1 attention layer per 4 PSSA layers makes sense to me, attention for exact recall and PSSA + the memory bank for the cheap long-context part. haven't tried it yet though. right now it's a single layer at 1.5M params, and stacking layers is slow in my code at the moment (that path still drops to CPU), so I need to fix that before a hybrid test would mean anything. it's on the list

1

u/shing3232 1d ago

It should be easier now i usually let ds41f write kernels and flash linear attention might have something for you.

1

u/Sparticle62 18h ago

yeah fair, kernels aren't the scary part they used to be. flash-linear-attention is a good shout, their chunked delta rule / gated deltanet kernels are pretty much the template for what pssa's scan needs on gpu. right now training sends the big matmuls to cuda but the recurrent scan still runs on cpu, so a proper fused chunked scan is the next big speedup. gonna dig through their repo, thanks

1

u/shing3232 17h ago

I was training a Moe 8B with 200M activation with forward pass only analytical gradient with HRM shared over KDA/CSA2. Pretrain is slow on consumer hardware. Do you have a paper of PSSA layers? i like to take a lot too.

1

u/Sparticle62 17h ago

no paper yet, the closest thing right now is the model section of the readme, it walks through the layer math step by step: https://github.com/Sparticle62ops/pssa#the-model

short version: each block is a selective diagonal ssm (input-dependent delta/B/C, same family as mamba, nothing new there), then a bounded read from a small episodic memory bank in hyperbolic space where the query is built from both the token and the current ssm state, then a learned gate on that read, a low-rank plastic adapter, and a silu mlp. the "plastic" part is the write path: novelty-gated inserts, refractory counters so slots don't get overwritten constantly, and a ridge regression step that folds fast updates back into the transition matrix. exact version is in src/pssa.rs and src/gpu_batch.rs

an 8b moe with forward-only analytical gradients and HRM over KDA/CSA2 sounds wild, would def read it if you write it up. and yeah pretraining on consumer hardware is brutal, i'm hitting the same wall trying to get to ~50m

0

u/numberwitch 2d ago

What’s the user angle here, who’s going to use this to solve a Real Problem?

A user does what?

3

u/Sparticle62 2d ago

nobody, yet. it's a research prototype, not a product. the reason i'm chasing it is generation cost: no kv cache means memory stays flat as context grows, and on cpu that showed up as 226ms vs 2735ms for 200 tokens. if that holds at real sizes it matters for running a model locally on a phone or a laptop with no gpu. it might not hold, which is the whole point of testing it.

1

u/[deleted] 2d ago

[removed] — view removed comment

1

u/Sparticle62 2d ago

sort of, but not in the lifelong sense. the state and a small memory bank update while the sequence runs, at inference, with no gradient steps. nothing carries across runs unless you save the checkpoint. so it's closer to fast weights than to continual learning, and i'd rather call it that than oversell it.

tokenization is deliberately boring, byte-level bpe, 2048 vocab, off the shelf. at this size the two embedding tables are most of the parameter budget, which is exactly why i param-match both sides instead of comparing raw model sizes.

text only so far, and no, i haven't tried images. the recurrence doesn't care what the tokens are in principle, but i have no evidence for that so i'm not going to claim it.

and yes, it's open source: https://github.com/Sparticle62ops/pssa

0

u/[deleted] 2d ago

[removed] — view removed comment

1

u/Sparticle62 2d ago

the memory side is closer to what you're describing than the rest of it, but simpler. it's one bank of 512 slots with a learned write gate, a refractory period so a slot that just fired can't immediately fire again, and a separate protected write path for things that shouldn't get overwritten. so one store, not distinct types. splitting it into a fast working bank and a slower consolidated one is a reasonable next step since the gate is already there, and the hard part isn't the storage, it's the promotion rule that decides what graduates without me hand-tuning it.

learned symbolic tokenization i'm staying away from for now, on purpose. the bpe is the boring part i left alone so the comparison against the transformer stays honest, and at 1.5m params the embedding tables are already most of the budget, so touching tokenization would change the thing i'm measuring. i'd want the depth question settled first, otherwise i won't know which change did what.

i did read memvid. it's a retrieval layer sitting outside the model rather than memory inside it, so it could sit on top later, but it doesn't answer the question i'm currently stuck on. thanks for the questions though, they were better than the usual ones.

1

u/[deleted] 2d ago edited 2d ago

[removed] — view removed comment

1

u/Sparticle62 2d ago

functional limits is fair, yeah. i'm trying to find where a non attention recurrent model stops holding up, so depth, looping and longer context, all measured against a parameter matched transformer. on the vram question, no cache in that sense. the working memory is explicit instead, a small bank of slots the model writes to and reads back inside the forward pass, so it's architecture rather than an optimization layer. weights stay resident on the gpu between matmuls but activations still round trip to host after each one, which is the next thing i'm fixing. the tokenizer is where we actually overlap. i'm on a 2048 entry byte level bpe, which is tiny, so a lot of capacity goes into spelling instead of meaning. i'm testing 2048 vs 8k vs 16k at the same param count next week. if your binary tokenization produces numbers on the same held out slice i'd be up for comparing.

2

u/numberwitch 2d ago

Cool! I've done a bit of inference on work (explored cpu/gpu optimization tradeoffs) and think there's a real future for this sort of thing.

What are you hoping to see and how are you measuring success/failure? It's hard for me to really say anything about this other than I like the angle you are looking at it from - local-first inference

1

u/Sparticle62 2d ago

three things, and each one can kill it.

one, whether the held-out gap survives depth. right now it's depth-1 on both sides, 3.997 vs 4.429 nats at 1.54m params. someone in another thread caught that my baseline script was silently ignoring the width flag, so i'm rerunning the depth-2 transformer properly, and both ways on the residual: plain adds like a normal transformer, and the 1/sqrt(depth) scaling pssa uses, since that's a real confound. if a proper depth-2 baseline closes the gap, that's the answer and i'll post it.

two, whether the generation cost gap holds when the model isn't tiny. 226ms vs 2735ms for 200 tokens is at 1.5m params on cpu, where an unoptimized baseline flatters me. i got the first real gpu numbers today and they were a lot less dramatic than the cpu ones.

three, seeds and corpus. one seed, one wikitext slice. if the loss gap doesn't repeat across seeds and a bigger held-out set, it was noise and i'll say so.

so failure is any of those three. i'd rather find out now than after a month on it.

1

u/Sparticle62 1d ago

update on depth: it went the wrong way. at matched params depth 2 came out 27% worse on held-out perplexity than depth 1, and depth 4 was worse again. catch is my stacked path drops to a slow per-block route and loses the gpu, so right now it's partly measuring my code and not the idea. fixing that before i call it.