Modele și platforme AI

Jamba: Noul model hibrid de limbaj Transformer-Mamba de la AI21 Labs

mm
Adaugă Unite.AI la sursele tale preferate pe Google

Modelele de limbaj au cunoscut o evoluție rapidă, arhitecturile bazate pe Transformer conducând schimbarea în procesarea limbajului natural. Cu toate acestea, pe măsură ce modelele se extind, provocările legate de gestionarea contextelor lungi, eficiența memoriei și debitul au devenit mai pronunțate.

AI21 Labs a introdus o nouă soluție cu Jamba, un model de limbaj mare de ultimă generație (LLM) care combină punctele forte ale arhitecturilor Transformer și Mamba într-un cadru hibrid. Acest articol detaliază arhitectura, performanța și aplicațiile potențiale ale lui Jamba.

Prezentare generală a lui Jamba

Jamba este un model de limbaj mare hibrid dezvoltat de AI21 Labs, care utilizează o combinație de straturi Transformer și Mamba, integrate cu un modul Mixture-of-Experts (MoE). Această arhitectură permite lui Jamba să echilibreze utilizarea memoriei, debitul și performanța, făcându-l un instrument puternic pentru o gamă largă de sarcini de procesare a limbajului natural. Modelul este proiectat să se încadreze într-un singur GPU de 80GB, oferind un debit ridicat și o amprentă mică de memorie, menținând în același timp o performanță de ultimă generație pe diverse benchmark-uri.

Arhitectura lui Jamba

Arhitectura lui Jamba este piatra de temelie a capacităților sale. Este construit pe un design hibrid inovator care intercalează straturi Transformer cu straturi Mamba, integrând module MoE pentru a îmbunătăți capacitatea modelului fără a crește semnificativ cerințele computaționale.

1. Straturi Transformer

Arhitectura Transformer a devenit standardul pentru modelele moderne de limbaj mare datorită capacității sale de a gestiona procesarea paralelă eficient și de a captura dependențele pe termen lung în text. Cu toate acestea, performanța sa este adesea limitată de cerințele ridicate de memorie și calcul, în special atunci când se procesează contexte lungi. Jamba abordează aceste limitări prin integrarea straturilor Mamba, pe care le vom explora în continuare.

2. Straturi Mamba

Mamba este un model de stare spațială (SSM) recent proiectat pentru a gestiona relațiile pe termen lung în secvențe mai eficient decât RNN-urile tradiționale sau chiar Transformer. Straturile Mamba sunt în special eficiente în reducerea amprentei de memorie asociate cu stocarea cache-ului cheie-valoare (KV) în Transformer. Prin intercalarea straturilor Mamba cu straturile Transformer, Jamba reduce utilizarea generală a memoriei, menținând în același timp o performanță ridicată, în special în sarcinile care necesită gestionarea contextelor lungi.

3. Module Mixture-of-Experts (MoE)

Modulul MoE din Jamba introduce o abordare flexibilă pentru scalarea capacității modelului. MoE permite modelului să crească numărul de parametri disponibili fără a crește proporțional parametrii activi în timpul inferenței. În Jamba, MoE este aplicat la unele dintre straturile MLP, mecanismul de rutare selectând cei mai buni experți pentru a fi activați pentru fiecare token. Această activare selectivă permite lui Jamba să mențină o eficiență ridicată în timp ce gestionează sarcini complexe.

Imaginea de mai jos demonstrează funcționalitatea unui cap de inducție într-un model hibrid de atenție-Mamba, o caracteristică cheie a lui Jamba. În acest exemplu, capul de atenție este responsabil pentru predicția etichetelor, cum ar fi “Pozitiv” sau “Negativ”, în sarcini de analiză a sentimentului. Cuvintele evidențiate ilustrează modul în care atenția modelului este puternic focalizată pe token-urile de etichetă din exemplele cu puține shot-uri, în special în momentul critic înainte de a prezice eticheta finală. Mecanismul de atenție joacă un rol crucial în capacitatea modelului de a efectua învățarea în context, unde modelul trebuie să inferă eticheta adecvată pe baza contextului dat și a exemplelor cu puține shot-uri.

Îmbunătățirile de performanță oferite prin integrarea MoE cu arhitectura hibridă de atenție-Mamba sunt evidențiate în Tabel. Prin utilizarea MoE, Jamba crește capacitatea sa fără a crește proporțional costurile computaționale. Acest lucru este evident în special în creșterea semnificativă a performanței pe diverse benchmark-uri, cum ar fi HellaSwag, WinoGrande și Întrebări Naturale (NQ). Modelul cu MoE nu numai că atinge o acuratețe mai ridicată (de exemplu, 66,0% la WinoGrande, comparativ cu 62,5% fără MoE), dar demonstrează și log-probabilități îmbunătățite pe diverse domenii (de exemplu, -0,534 la C4).

Caracteristici arhitecturale cheie

  • Compoziția stratului: Arhitectura lui Jamba constă în blocuri care combină straturi Mamba și Transformer într-un raport specific (de exemplu, 1:7, ceea ce înseamnă un strat Transformer pentru fiecare șapte straturi Mamba). Acest raport este reglat pentru o performanță și eficiență optimă.
  • Integrarea MoE: Straturile MoE sunt aplicate la fiecare câteva straturi, cu 16 experți disponibili și cei doi experți superiori activați pe token. Această configurație permite lui Jamba să se extindă eficient, gestionând compromisurile dintre utilizarea memoriei și eficiența computațională.
  • Normalizare și stabilitate: Pentru a asigura stabilitatea în timpul antrenamentului, Jamba incorporează RMSNorm în straturile Mamba, ceea ce ajută la mitigarea problemelor, cum ar fi spike-urile de activare mari care pot apărea la scară.

Performanța și benchmarking-ul lui Jamba

Jamba a fost testat riguros împotriva unei game largi de benchmark-uri, demonstrând o performanță competitivă pe tot parcursul. Următoarele secțiuni evidențiază unele dintre principalele benchmark-uri în care Jamba a excelat, demonstrându-și punctele forte atât în sarcinile generale de procesare a limbajului natural, cât și în scenariile cu context lung.

1. Benchmark-uri comune de procesare a limbajului natural

Jamba a fost evaluat pe diverse benchmark-uri academice, incluzând:

  • HellaSwag (10-shot): O sarcină de raționament comun în care Jamba a atins un scor de performanță de 87,1%, depășind multe modele concurente.
  • WinoGrande (5-shot): O altă sarcină de raționament în care Jamba a obținut un scor de 82,5%, demonstrând din nou capacitatea sa de a gestiona raționamente lingvistice complexe.
  • ARC-Challenge (25-shot): Jamba a demonstrat o performanță puternică cu un scor de 64,4%, reflectând capacitatea sa de a gestiona întrebări multiple dificile.

În benchmark-uri agregate, cum ar fi MMLU (5-shot), Jamba a obținut un scor de 67,4%, indicând robustețea sa pe sarcini diverse.

2. Evaluări de context lung

Una dintre caracteristicile deosebite ale lui Jamba este capacitatea sa de a gestiona contexte extrem de lungi. Modelul suportă o lungime de context de până la 256K de tokeni, cea mai lungă dintre modelele disponibile public. Această capacitate a fost testată utilizând benchmark-ul Needle-in-a-Haystack, unde Jamba a demonstrat o acuratețe remarcabilă de recuperare pe diverse lungimi de context, incluzând până la 256K de tokeni.

3. Debit și eficiență

Arhitectura hibridă a lui Jamba îmbunătățește semnificativ debitul, în special cu secvențe lungi.

În teste care compară debitul (tokeni pe secundă) între diverse modele, Jamba a depășit constant concurenții săi, în special în scenarii care implică dimensiuni mari de lot și contexte lungi. De exemplu, cu un context de 128K de tokeni, Jamba a atins un debit de trei ori mai mare decât Mixtral, un model comparabil.

Folosirea lui Jamba: Python

Pentru dezvoltatori și cercetători dornici să experimenteze cu Jamba, AI21 Labs a pus modelul la dispoziție pe platforme precum Hugging Face, făcându-l accesibil pentru o gamă largă de aplicații. Următorul fragment de cod demonstrează cum să încărcați și să generați text utilizând Jamba:


<p>from transformers import AutoModelForCausalLM, AutoTokenizer</p>

<p>model = AutoModelForCausalLM.from_pretrained(&quot;ai21labs/Jamba-v0.1&quot;)
tokenizer = AutoTokenizer.from_pretrained(&quot;ai21labs/Jamba-v0.1&quot;)</p>

<p>input_ids = tokenizer(&quot;În Super Bowl LVIII recent,&quot;, return_tensors=&#039;pt&#039;).to(model.device)[&quot;input_ids&quot;]</p>

<p>outputs = model.generate(input_ids, max_new_tokens=216)</p>

print(tokenizer.batch_decode(outputs))

Acest script simplu încarcă modelul Jamba și tokenizer-ul, generează text pe baza unei intrări date și imprimați ieșirea generată.

Reglarea fină a lui Jamba

Jamba este proiectat ca un model de bază, ceea ce înseamnă că poate fi reglat pentru sarcini sau aplicații specifice. Reglarea fină permite utilizatorilor să adapteze modelul la domenii de nișă, îmbunătățind performanța pe sarcini specializate. Următorul exemplu arată cum să reglați fin Jamba utilizând biblioteca PEFT:

import torch
from datasets import load_dataset
from trl import SFTTrainer, SFTConfig
from peft import LoraConfig
from transformers import AutoTokenizer, AutoModelForCausalLM, TrainingArguments

<p>tokenizer = AutoTokenizer.from_pretrained(&quot;ai21labs/Jamba-v0.1&quot;)
model = AutoModelForCausalLM.from_pretrained(
&quot;ai21labs/Jamba-v0.1&quot;, device_map=&#039;auto&#039;, torch_dtype=torch.bfloat16)</p>

<p>lora_config = LoraConfig(r=8,
target_modules=[
&quot;embed_tokens&quot;,&quot;x_proj&quot;, &quot;in_proj&quot;, &quot;out_proj&quot;, # mamba
&quot;gate_proj&quot;, &quot;up_proj&quot;, &quot;down_proj&quot;, # mlp
&quot;q_proj&quot;, &quot;k_proj&quot;, &quot;v_proj&quot;
# atenție],
task_type=&quot;CAUSAL_LM&quot;, bias=&quot;none&quot;)</p>

<p>dataset = load_dataset(&quot;Abirate/english_quotes&quot;, split=&quot;train&quot;)
training_args = SFTConfig(output_dir=&quot;./results&quot;,
num_train_epochs=2,
per_device_train_batch_size=4,
logging_dir=&#039;./logs&#039;,
logging_steps=10, learning_rate=1e-5, dataset_text_field=&quot;quote&quot;)
trainer = SFTTrainer(model=model, tokenizer=tokenizer, args=training_args,
peft_config=lora_config, train_dataset=dataset,
)
trainer.train()

Acest fragment de cod reglează fin Jamba pe un set de date de citate în engleză, ajustând parametrii modelului pentru a se potrivi mai bine sarcinii specifice de generare de text într-un domeniu specializat.

Implementarea și integrarea lui Jamba

AI21 Labs a făcut familia Jamba larg accesibilă prin diverse platforme și opțiuni de implementare:

  1. Platforme cloud:
    • Disponibil pe principalele furnizori de cloud, incluzând Google Cloud Vertex AI, Microsoft Azure și NVIDIA NIM (NVDA ).
    • Urmează să fie disponibil pe Amazon Bedrock, Databricks Marketplace și Snowflake Cortex.
  2. Cadrul de dezvoltare AI:
    • Integrare cu cadre populare, cum ar fi LangChain și LlamaIndex (urmează să fie lansat).
  3. AI21 Studio:
    • Acces direct prin platforma de dezvoltare a AI21.
  4. Hugging Face:
    • Modele disponibile pentru descărcare și experimentare.
  5. Implementare locală:
    • Opțiuni pentru implementare privată, pe site, pentru organizații cu nevoi specifice de securitate sau conformitate.
  6. Soluții personalizate:
    • AI21 oferă servicii de personalizare și reglare fină a modelului pentru clienții enterprise.

Caracteristici prietenoase pentru dezvoltatori

Modelele Jamba vin cu mai multe capacități încorporate care le fac deosebit de atractive pentru dezvoltatori:

  1. Apelarea funcțiilor: Integrați ușor instrumente și API-uri externe în fluxurile de lucru AI.
  2. IEșire JSON structurată: Generați structuri de date curate și parsabile direct din intrări de limbaj natural.
  3. Digestia obiectului document: Procesați și înțelegeți eficient structuri de document complexe.
  4. Optimizări RAG: Caracteristici încorporate pentru a îmbunătăți pipe-line-urile de generare augmentată de recuperare.

Aceste caracteristici, combinate cu fereastra lungă de context și procesarea eficientă a modelului, fac din Jamba un instrument versatil pentru o gamă largă de scenarii de dezvoltare.

Considerații etice și AI responsabil

În timp ce capacitățile lui Jamba sunt impresionante, este crucial să abordăm utilizarea sa cu o mentalitate de AI responsabil. AI21 Labs subliniază mai multe puncte importante:

  1. Natura modelului de bază: Modelele Jamba 1.5 sunt modele de bază preantrenate fără aliniere sau reglare specifică a instrucțiunilor.
  2. Lipsa mecanismelor de protecție încorporate: Modelele nu au mecanisme de moderare încorporate.
  3. Implementare atentă: Adaptarea și mecanismele de protecție suplimentare ar trebui implementate înainte de a utiliza Jamba în medii de producție sau cu utilizatori finali.
  4. Confidențialitatea datelor: Atunci când se utilizează implementări bazate pe cloud, fiți conștienți de manipularea și cerințele de conformitate a datelor.
  5. Conștientizarea prejudecăților: Ca și toate modelele de limbaj mare, Jamba poate reflecta prejudecățile prezente în datele de antrenament. Utilizatorii ar trebui să fie conștienți de acest lucru și să implementeze măsuri de atenuare adecvate.

Prin a ține cont de acești factori, dezvoltatorii și organizațiile pot utiliza capacitățile lui Jamba într-un mod responsabil și etic.

Un nou capitol în dezvoltarea AI?

Introducerea familiei Jamba de către AI21 Labs marchează un punct de referință semnificativ în evoluția modelelor de limbaj mare. Prin combinarea punctelor forte ale arhitecturilor Transformer și a modelelor de stare spațială, integrând tehnici de mixture of experts și împingând limitele lungimii contextului și a vitezei de procesare, Jamba deschide noi posibilități pentru aplicații AI în diverse industrii.

Pe măsură ce comunitatea AI continuă să exploreze și să construiască pe această arhitectură inovatoare, putem aștepta să vedem progrese suplimentare în eficiența modelului, înțelegerea contextului lung și implementarea practică a AI. Familia Jamba reprezintă nu numai un nou set de modele, ci și o posibilă schimbare în abordarea proiectării și implementării sistemelor AI de scară largă.

Am petrecut ultimii cinci ani scufundându-mă în lumea fascinantă a Machine Learning și Deep Learning. Pasinea și expertiza mea m-au condus să contribui la peste 50 de proiecte diverse de inginerie software, cu un focus deosebit pe AI/ML. Curiozitatea mea în continuare m-a atras și spre Natural Language Processing, un domeniu pe care sunt dornic să îl explorez mai departe.