Implementação do Mamba em um único arquivo PyTorch
(github.com/johnma2006)- mamba-minimal é um projeto que implementa o Mamba de forma simples e mínima em um único arquivo PyTorch
- O objetivo é produzir a mesma saída numérica da implementação oficial no forward/backward pass
- O código foi simplificado e estruturado em um formato comentado para facilitar a leitura
- Não inclui as otimizações centrais da implementação oficial, portanto não oferece velocidade, nem inclui inicialização adequada de parâmetros
- A demo executa um exemplo de conclusão de prompt usando
state-spaces/mamba-370me o tokenizadorEleutherAI/gpt-neox-20b
Visão geral do projeto
- mamba-minimal é uma implementação simples e mínima do Mamba em um único arquivo PyTorch
- O objetivo é mostrar o mesmo comportamento da implementação oficial com um código mais fácil de ler
- Principais características:
- Saída numérica equivalente à implementação oficial no forward pass e no backward pass
- Código simplificado
- Implementação comentada e fácil de ler
O que não está incluído
- Velocidade não é o objetivo
- A implementação oficial é fortemente otimizada
- Essas otimizações fazem parte da principal contribuição do artigo do Mamba
- Esta implementação mantém a maior parte do código simplificada em favor da legibilidade
- Inicialização adequada de parâmetros não está incluída
- Isso é apresentado como algo que pode ser adicionado sem sacrificar a legibilidade
Exemplo de uso da demo
- É possível ver um exemplo de conclusão de prompt em
demo.ipynb - O exemplo usa
model.Mambae oAutoTokenizerdo Hugging Facetransformers - Modelo e tokenizador usados:
state-spaces/mamba-370mEleutherAI/gpt-neox-20b
- O prompt de exemplo é
Mamba is the, e o resultado gerado inclui uma frase descrevendo a mamba como uma cobra venenosa
Referências
- A arquitetura Mamba é apresentada em Mamba: Linear-Time Sequence Modeling with Selective State Spaces
- Os autores do artigo são Albert Gu e Tri Dao
- A implementação oficial está no repositório
state-spaces/mamba
1 comentários
Opiniões do Hacker News
Há algum tempo, com um colega, criei uma biblioteca que separava a maior parte do código de modelos compartilhado; usando-a, é possível implementar muitos modelos em cerca de 100 linhas, excluindo imports em Python e comentários
BERT: https://github.com/explosion/curated-transformers/blob/main/...
Llama 1/2: https://github.com/explosion/curated-transformers/blob/main/...
MPT: https://github.com/explosion/curated-transformers/blob/main/...
Também oferece suporte a recursos como TorchScript JIT e PyTorch flash attention
O xformers também aborda problemas parecidos, mas é mais focado em fornecer módulos Transformer de alto desempenho usando Triton. Porém, não foi fácil usar apenas componentes específicos da biblioteca, e erros em tempo de execução continuavam aparecendo, então deixei de lado por enquanto. Como estou criando algo baseado na arquitetura BERT, vou usá-la como referência
Pretendo usá-la no meu próximo projeto
O código original do Mamba tem muitas otimizações de velocidade e outros elementos, então não é fácil entendê-lo de imediato; esta implementação deve ajudar no aprendizado
Ao fazer inferência token por token, tudo fica muito mais simples. Também tenho uma implementação própria de inferência do Mamba: https://github.com/rbitr/llm.f90/tree/master/ssm
Sei que ele é a base de código de computação científica há muito validado e que muitas vezes é usado por trás de bibliotecas como PyTorch ou Numpy, mas não é uma linguagem popular hoje em dia. Gostaria de saber o motivo da escolha
Há partes sobre o Mamba que eu gostaria que fossem explicadas de um jeito que até quem não é pesquisador de machine learning consiga entender:
Ri da frase “a mamba é a cobra venenosa mais longa do mundo, com comprimento estimado acima de 150 m”
Ainda assim, é realmente excelente, e achei bom que tenha referenciado o artigo no arXiv, porque pessoas como eu, que consomem textos assim em vez de interpretar o artigo diretamente, conseguem dar uma espiada por dentro
Seria engraçado se depois tivessem que publicar uma correção para essa frase
Eu esperava que o núcleo do algoritmo fosse o parallel prefix scan. Acho que esse é o ponto do 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)Pode ser uma pergunta boba, mas fico curioso sobre a dificuldade de treinar o modelo Mamba disponível no Hugging Face
O maior modelo parece ser de 2,8B; quantas GPUs seriam necessárias e quanto tempo levaria para treiná-lo em um dataset como The Pile?
A inferência também parece rodar 3 a 5 vezes mais rápido, usando apenas metade da RAM
Tentei destrinchar a versão CUDA oficial, mas depois que a primeira tentativa falhou acabei não mexendo mais; esta implementação parece muito melhor
Mais uma implementação em PyTorch em arquivo único, realmente excelente. Espero que trabalhos anteriores como hlb-CIFAR10 e projetos relacionados, além de influências anteriores como minGPT e DawnBench, tenham ajudado a impulsionar nem que seja um pouco esse formato simples de arquivo único
Esse tipo de trabalho é importante para pesquisa eficiente em machine learning, e talvez seja uma das coisas mais importantes que se pode fazer agora pela área
A pesquisa avança no ritmo da inovação, a inovação acelera na proporção inversa ao tempo de execução dos experimentos, e isso claramente tem relação com a complexidade de Kolmogorov do código para fins de pesquisa ou hacking simples
Não dá para enfatizar o bastante o quanto ferramentas assim são importantes para a pesquisa e o quanto, pessoalmente, elas aceleraram meu processo de descoberta de conhecimento. A capacidade de rascunhar ideias rapidamente em minutos e receber resultados imediatos com alta relação sinal-ruído se tornou essencial para o progresso da pesquisa
Vejo destilação de conhecimento e MDL(https://en.wikipedia.org/wiki/Minimum_description_length) como muito importantes para reverter os enfeites desnecessários, a bagunça e a competição excessivamente densa por temas de baixo valor para “não ficar para trás” que o processo atual de submissão e revisão de artigos parece incentivar
Recentemente, para evitar esse problema e caminhar para uma solução de escala um pouco melhor, comecei a distribuir código como “esboços de código”: gists curtos, autocontidos, de um único arquivo. Isso reduz o tempo de desenvolvimento e permite entregar diretamente às pessoas código funcional, bruto e sem polimento, que contém o conceito. Até agora parece estar funcionando bem, e quero continuar
Gostaria de ver mais código assim. Se são pesquisadores que treinam dados em larga escala, também deveriam ser eficientes em dados na forma como disseminam informação
Provavelmente a humanidade nunca desenvolveu algo com complexidade considerável tão rápido assim
Um lugar onde se vê uma velocidade parecida talvez seja a SpaceX, que também lançou dois foguetes de ponta este ano. Fico curioso para saber o que virá em 2024
x_projnão tem bias, parece que seria possível combinar os pesos de x_proj e dt_projSe houver exigências de ajuste de pesos, talvez isso possa ser feito simplesmente em tempo de execução, e no fim um único kernel com bias provavelmente será mais rápido. Não tenho certeza
Fico curioso se houve uma discussão sobre o artigo original. Acho que perdi, mas é bem interessante
Não entendi bem a parte que diz: “por falta de uma implementação eficiente, há falta de memória ou exigências computacionais irrealistas, então faltam os resultados completos para comprimento de contexto de 8k dos baselines RWKV e RetNet, modelos recorrentes fortes anteriores que também podem ser interpretados como SSMs”
RetNet não usa muita memória e, usando uma implementação de forward em chunks, o uso de VRAM fica limitado pelo tamanho do chunk. Essa é a essência de testar comprimento de contexto
Fico curioso se alguém testou o modelo Mamba original. Qual seria a velocidade de treinamento em comparação com o RetNet no modo de forward paralelo?
https://openreview.net/forum?id=AL1fq05o7H
Implementações que reduzem algo complexo ao essencial são sempre boas