10 ms·
Implementation of Mamba in one file of PyTorch
- andy99 3y agoThe original mamba code has a lot of speed optimizations and other stuff that make it difficult to immediately get so this will help with learning. I can't help but also plug my own Mamba inference implementation. https://github.com/rbitr/llm.f90/tree/master/ssm https://github.com/rbitr/llm.f90/tree/master/ssm For inference one token at a time everything simplifies considerably.
- cs702 3y agoFortran! If you don't mind me asking, why Fortran? I know it underpins a lot of time-tested scientific code, often wrapped by libraries like PyTorch and Numpy, but Fortran isn't exactly a popular language nowadays. What's your rationale for using it?
- andy99 3y agoTdlr, Fortran is low level-ish, compiled, but otherwise almost identical to numpy syntax wise. It supports all the common array and matrix operations and it doesn't need memory and pointer management the way C does. But it still compiles down to something very fast, you can link in BLAS and GPU libraries, supports easy parallelism... When I compare with e.g. Karpathy's llama2.c, I think Fortran is easy to work with implementing basic transformer inference because of how it handles arrays. The downside is that while there are efforts to modernize it, I find it more cumbersome for non-numerical stuff, particularly strings. But I think for the actual linear algebra implementation, it can't be beat. I should add, I know it's a bit of an uphill battle, I expect fewer people will use code that I write in Fortran vs basically anything else. But I'm hoping to pull some people in and get a critical mass of interest because I think it has a lot of promise. That's actually one of the reasons I wanted to get a Mamba implementation quickly (though now that there's a basic python one I think I'll have lost some potential users to it :)
- cs702 3y agoThanks for the thoughtful response. Unfortunately, I too think it will be a bit of an uphill battle for you. If you haven't already, take a look at Mojo and Julia. Both offer many of the benefits of Fortran, but unlike it, they are seeing growing adoption.
- andy99 3y agoAn uphill battle is fine
- gnaritas99 3y ago[dead]
- cztomsik 3y agoI only heard good things about fortran :)
- cs702 3y agoThis looks really nice. Thank you for sharing it on HN! In case you didn't know, you can parallelize the slow Python loop in selective_scan that computes all the x's: x = torch.zeros((b, d_in, n)) for i in range(l): x = deltaA[:, :, i] * x + deltaB_u[:, :, i] ⋮ with only two calls to the PyTorch API. See the examples here: https://github.com/glassroom/heinsen_sequence/blob/main/README.md https://github.com/glassroom/heinsen_sequence/blob/main/READ... .[a] You can then compute all the y's with one einsum, instead of l sequential einsums. --- [a] Previous discussion on HN: https://news.ycombinator.com/item?id=38556669 https://news.ycombinator.com/item?id=38556669
- deleted 3y ago[deleted]
- make3 3y agoOP's code is much easier to understand, though, which is the main (only) purpose of their code
- cs702 3y agoCan't argue with that! :-) For what it's worth, you can keep both, and make parallel vs sequential execution an option, with a boolean flag. You can also leave the sequential code as a comment explaining what the parallel code does. Or, if slow execution doesn't bother you, leave it as is.
- bradfitz 3y agoYou're replying to somebody who was arguing for readability being its virtue and you're proposing ... adding options and alternate code paths? :)
- anytime5704 3y agoVia a boolean parameter, no less.
- 3y ago
- boredumb 3y ago"Mamba is the world's longest venomous snake with an estimated length of over 150 m" Had a laugh at that. Really great stuff though, it was nice to have referencing to the arxiv paper so someone like me who generally consumes these things instead of translating them from papers could sort of peak behind the curtains.
- visarga 3y agoMamba has a great name ... [S]elective [S]tructured [S]tate [S]pace [S]equence models.. makes sSSSS, like a snake
- behnamoh 3y agoIf only the "mamba" name were not ugly.
- rdedev 3y agoWait I thought that was the king cobra? The longest venomous snake ? At least that was what a simple Google search showed me. Would be funny if they had to issue a correction for that sentence later on
- y42 3y agoslightly OT: I really struggle with dozens and dozens of vocabulary that is being used in the field of machine learning and especially AI. I'm not a beginner at all, but I wonder if there is a comprehensive guide for all those terms that not necessarily explains the technology behind them in detail, but shows their position and relation to each other. like some kind of landscape. "everyone" seems to know Mamba. I never heard of Mamba. There are constantly new kind of llm popping up, talking about stuff that seems to be obvious. So, is there some kind of resource like that, not aiming at beginners, but experienced users, coming from other fields of IT?
- bananaflag 3y agoI knew about Mamba from r/singularity and following AI researchers on Twitter. I don't work in AI at all (and don't plan to), but it's fun to know about stuff a little before they become mainstream.
- sevagh 3y ago>"everyone" seems to know Mamba. I never heard of Mamba Only the "everybody who knows what mamba is" are the ones upvoting and commenting. Think of all the people who ignore it. For me, Mamba is the faster version of Conda [1], and that's why I clicked on the article. https://github.com/mamba-org/mamba https://github.com/mamba-org/mamba
- 3-cheese-sundae 3y agoAh yes, Conda, definitely something else I've heard of.
- NavinF 3y agoConda has been around for a decade and it used to be the primary package manager for everything related to numpy/scipy. Most ML and data science people have heard of it even if they haven't used it.
- sevagh 3y agoConda is the latest LLM cli frontend that's a MOE of Mistral 7B, LLama 17B, Falcon 32C, and the Yamaha YZ50 quad bike.
- swyx 3y agothings I'd like a non-ML-researcher explanation of about Mamba: 1. what is the overall insight of state space models beyond transformers? (i know this is somewhat covered in the paper but still a bit inaccessible) 2. what was the incremental innovation/result that is making Mamba more successful/interesting than its predecessors? (S4, H3, Monarch etc) 3. what are the implications beyond subquadratic scaling of context? say if i don't really care about context length > 100k tokens. what other benefits are there - for example, is Mamba potentially more compute-efficient to train for a similar size of model/dataset? just offering 3 prompts for knowledgeable people to drop some alpha
- logicchains 3y agoFor 2, Mamba makes some A B C weights that in S4 are time invariant become functions of the input, which makes it more powerful.
- pk-protect-ai 3y ago> is Mamba potentially more compute-efficient to train for a similar size of model/dataset? I would like to understand it too as well ... Here is the citation from original paper: "Computation. After the parameters have been transformed from (∆, A, B, C) ↦ (A, B, C), the model can be computed in two ways, either as a linear recurrence (2) or a global convolution (3). Commonly, the model uses the convolutional mode (3) for efficient parallelizable training (where the whole input sequence is seen ahead of time), and switched into recurrent mode (2) for efficient autoregressive inference (where the inputs are seen one timestep at a time)." So the training is parallelizable, like in RetNet with parallel forward mode. By default inference is done in the recurrent mode, to have a longest possible context. No chunking available, so it is difficult for me to say how much RAM and VRAM it will consume during the inference ...
- pk-protect-ai 3y agoI did some minimal testing, mamba uses about 60% of VRAM in comparison to RetNet (parallel forward mode) with the model of the same size and the vocabulary of same size during inference.
- mcemilg 3y agoLooks wonderful. But I would like to add this, I hate einops, it doesn't make it simple to read unfortunately.
- andy99 3y agoI re-implemented Mamba myself and this was the first time I had ever worked with einops/einsum. I'm 50/50 on them after this. I found them relatively easy to look at and understand the intent (possibly more so than other representations), but talking extra time to transforms into other primitives (loops, multiplication, etc). I belive using torch.einsum is generally well optimized as well compared to naively looping. All said, I don't know if I'd use it myself working from scratch but it's interesting to know and if I was working in python I might try comparing the speed of einops/sum vs other ways.
- sjkoelle 3y agodisagree
- microtonal 3y agoNice! For what it is worth, a colleague and I made a library a while ago that factors out most shared model code, with which many models can be implemented in about 100 lines (excluding Python import ceremony and comments). E.g.: BERT: https://github.com/explosion/curated-transformers/blob/main/curated_transformers/models/bert/encoder.py https://github.com/explosion/curated-transformers/blob/main/... Llama 1/2: https://github.com/explosion/curated-transformers/blob/main/curated_transformers/models/llama/decoder.py https://github.com/explosion/curated-transformers/blob/main/... MPT: https://github.com/explosion/curated-transformers/blob/main/curated_transformers/models/mpt/decoder.py https://github.com/explosion/curated-transformers/blob/main/... With various stuff enabled, including support for TorchScript JIT, PyTorch flash attention, etc.
- rdedev 3y agoNice. I will definitely be taking a look at this. Have you looked at the xformers library ? They are looking at the same problem as you but their focus is more on providing performant transformer modules using triton. Using specific components from the library though is not as simple. I kept running into runtime errors so I've kept it aside for now. I am building something based on the Bert architecture so I will give this a look. Thanks for all the work!
- microtonal 3y agoI would've loved to look at xFormers, but I avoided looking at other implementations to make sure that ours is a clean room implementation. Curated Transformers started as a very small library just for spaCy (spaCy 3.7 transformer pipelines use Curated Transformers) with just the older encoder models (BERT, RoBERTa, etc.). spaCy used Hugging Face Transformers prior for the provided transformer models, but we wanted something where we could easily hook into different parts of the model (e.g. for distillation). After the functionality needed for spaCy was done, Matt @ Explosion encouraged us to extend it into a more general PyTorch library that would also support decoder architectures, generation, etc.
- deepsquirrelnet 3y agoI’m floored by this library. Great concept. I’ve never been a fan of HFs implementation. This is a beautiful API at exactly the right level of abstraction. I’ll give this a try on my next project.
- pk-protect-ai 3y agoIs there an original paper discussion? I seem to have missed it. It's quite interesting. I didn't catch on to this part: "We note that full results on context length 8k are missing for the RWKV and RetNet baselines, prior strong recurrent models that can also be interpreted as SSMs, due to a lack of efficient implementation leading to out-of-memory or unrealistic computation requirements." RetNet doesn't really consume much memory, and with the chunkwise forward implementation, it restricts the VRAM usage to the chunk size. This is the part to test the context length. Has anyone done some tests on the original Mamba model? How fast is the training on this one in comparison with RetNet in parallel forward mode?
- error9348 3y agohttps://news.ycombinator.com/item?id=38522428 https://news.ycombinator.com/item?id=38522428 https://openreview.net/forum?id=AL1fq05o7H https://openreview.net/forum?id=AL1fq05o7H
- MacsHeadroom 3y agoFaster training, much faster inference, and about half the VRAM usage during inference.
- allanrbo 3y agoLove it when complex things are distilled down to just the essentials!
- bobse 3y ago[dead]
- dheera 3y agoI love one file implementations. I hate all these implementations with preprocess_utils.py that imports stuff from model.py that imports stuff again from preprocess_utils.py that imports stuff from ...
- fassssst 3y agoFeels like a useful preprocessor script: turn this repo into a single file
- squigz 3y agoIs number of files in a project a meaningful metric...?
- Reubend 3y agoYes, it is here. This is an implementation designed for education: the main purpose here is to understand the model architecture in a practical sense. So lines of code and number of files are both meaningful. This is 1 short Python file, which makes it a lot easier to understand than a full optimized implementation.
- DarmokJalad1701 3y agoThanks for this. I took a stab at unraveling the official CUDA version and never really got around to it after my initial attempt failed. This seems a lot nicer.
- tysam_and 3y agoOh my gosh, another one-file PyTorch implementation. This is fantastic. I'd like to hope that some of my previous work (hlb-CIFAR10 and related projects, along with other influences before it like minGPT, DawnBench, etc.) has been able to help push the 'simple, single-file, reduced-complexity' format forward a bit. I personally think that this kind of work is critical to efficient ML research, and that is possibly one of the most important things that we can do for the field today. Research progresses at the speed of innovation, which progresses with the inverse of experiment runtime, which is definitely and absolutely related to the underlying Kolmogorov Complexity of the code w.r.t. a research/simple-hackery-focused objective. I really cannot stress enough how important to research tools like this are and how much they've sped up the knowledge discovery process for me personally. Being able to quickly sketch out ideas, often in minutes, and get immediate, high-snr results back has become an indispensable part of my research progress. While we seem to really good at some of the specifics of some of the detailsresearch, and somehow have extremely information-efficient training processes, we have not applied the same logic seemingly on the whole to the entire research field! Knowledge distillation and/or the MDL (https://en.wikipedia.org/wiki/Minimum_description_length https://en.wikipedia.org/wiki/Minimum_description_length) are excessively important I think to reversing a lot of the constant fluff, cruft, and overly dense thrash-and-hope-you-don't-get-scooped-by-other-researchers-on-marginal-value-topics trend that I think has largely been encouraged by the current paper submission/review/etc process. I've been wanting to try to get around this and move a bit more towards a slightly better scaling solution recently. One of these things is that I've started distributing my code in 1-file, self-contained, short rough gists as 'code sketches', which shortens dev time and gets rough, unpolished, working code for a concept in people's hands. It seems to work pretty well so far, I hope to continue doing it! <3 :')))) In any case, this is extremely exciting stuff, and everyone -- please! More code like this! We're researchers on learning data in a largely-scaled way, let's be data-efficient in how we disseminate information as well! It's a dream come true to see a lot more of this stuff coming down the pipeline, fantastic work and keep it coming! <3 :')))) Woop woop woop!!!! Excellent stuff. <3 :'))))
- tysam_and 3y agoMinor potential performance benefit -- it looks like you might be able to fuse the x_proj and dt_proj weights here as x_proj has no bias. This is a thing that's possibly doable simply at runtime if there's any weight-fiddling reqs, I'm guessing the single kernel + bias will still run faster in the end (not sure though! <3 :')))) )
- jdeaton 3y agoVery cool ive read this line of paper originating from hippo, s4, hyena, mamba etc but can someone please explain how this isnt just an RNN/LSTM variant??
- rhaps0dy 3y agoIts latent space transition is linear, instead of nonlinear, so there's a more parallelizable algorithm for advancing time in it. This makes it much more efficient to train and do inference with in GPUs. The way it keeps all the representation power of LSTMs is by having the transition vary with the input (but still be linear).
- jdeaton 3y agoThanks thats helpful. One place where the parallelizability of this method falls short of the transformer is not being able to pack multiple varying length examples into the same array during training with block diagonal attention pattern. If I understand correctly thats not possible with this architecture and its an important practical concern in large scale transformer training.
- epaulson 3y agoThis is a dumb question but how hard is it to train the mamba models that are on huggingface? It looks like the largest one is 2.8b - how many GPUs for how long do you need to train that up using a dataset like The Pile?
- MacsHeadroom 3y agoThat's a great question and I would like to know too. It looks like the answer is substantially faster than an equally sized Transformer, and the end result will score better than a Transformer on basically every benchmark. Also it will do inference 3-5x faster in half the RAM.
- uejfiweun 3y agoHow long does it generally take between model architectures like Mamba being proposed and the use of these architectures in SotA mega models like GPT or Gemini? IIUC Mamba basically eliminates restrictions on context length which would be awesome to see in the super-mega high performance models.
- brcmthrowaway 3y agoGPT-5 would have this enhancement
- MacsHeadroom 3y agoGPT-5 will not, because the T in GPT stands for Transformer and Mamba/SSMs/S6 are not Transformers. But I would bet that we see a SOTA S6 LLM from Meta by this Spring.
- brcmthrowaway 3y agoS6?
- marmaduke 3y agoHm I'd take a stab at a Jax version based on this. Thanks
- iskander 3y agoI expected the core of the algorithm to be a parallel prefix scan though (isn't that the point of Mamba?): for i in range(l): x = deltaA[:, :, i] \* x + deltaB_u[:, :, i] y = einsum(x, C[:, i, :], 'b d_in n , b n -> b d_in') ys.append(y)
- ekiauhce 3y agoIf a variable contains batch size, then name it accordingly — batch_size. And no glossary needed, KISS https://github.com/johnma2006/mamba-minimal/blob/82efa90919c3b5066674216f3edcebb3414a7b8f/model.py#L11 https://github.com/johnma2006/mamba-minimal/blob/82efa90919c...
- deleted 3y ago[deleted]
- DarmokJalad1701 3y agoI think the glossary is defining variable names as given in the paper. I found this confusing when I originally read the paper as the authors assume that the reader knows what B, L, D and N stand for. I had to use explainpaper to figure it out.
- Nimitz14 3y agoVery nice. Love the glossary!
- danielhanchen 3y agoCool share!
- Kim898 3y ago[dead]
- pama 3y agoSee also this resource if you’re interested in these models: https://news.ycombinator.com/item?id=38719675 https://news.ycombinator.com/item?id=38719675