AIモデルとプラットフォーム

Jamba:AI21 Labsの新しいハイブリッドTransformer-Mamba言語モデル

mm
Unite.AI を Google の優先ソースに追加

言語モデルは、Transformerベースのアーキテクチャが自然言語処理を牽引する中で急速に進化しています。しかし、モデルが拡大するにつれて、長いコンテキストの処理、メモリ効率、スループットの課題がより顕著になりました。

AI21 Labsは、TransformerとMambaアーキテクチャの強みをハイブリッドフレームワークで組み合わせた、最先端の大規模言語モデル(LLM)であるJambaを発表しました。この記事では、Jambaのアーキテクチャ、パフォーマンス、潜在的な応用について説明します。

Jambaの概要

Jambaは、AI21 Labsによって開発されたハイブリッド大規模言語モデルで、Transformer層とMamba層を組み合わせ、Mixture-of-Experts(MoE)モジュールを統合しています。このアーキテクチャにより、Jambaはメモリ使用量、スループット、パフォーマンスのバランスをとることができ、幅広いNLPタスクで強力なツールとなります。モデルは、単一の80GB GPU内に収まるように設計されており、高いスループットと小さなメモリフットプリントを提供しながら、さまざまなベンチマークで最先端のパフォーマンスを維持します。

Jambaのアーキテクチャ

Jambaのアーキテクチャは、その機能の根幹です。Transformer層とMamba層を交互に配置し、MoEモジュールを組み込むことで、計算要求を大幅に増やさずにモデルの容量を高めることができます。

1. Transformer層

Transformerアーキテクチャは、並列処理を効率的に処理し、テキスト内の長距離の依存関係を捉える能力があるため、近代的なLLMの標準となりました。しかし、そのパフォーマンスは、特に長いコンテキストを処理する場合に、高いメモリと計算要求によって制限されることがよくあります。Jambaは、Mamba層を統合することでこれらの制限に対処します。

2. Mamba層

Mambaは、シーケンス内の長距離の関係を効率的に処理するために設計された最近の状態空間モデル(SSM)です。Mamba層は、Transformerのキー値(KV)キャッシュのメモリフットプリントを削減するのに特に効果的です。Transformer層とMamba層を交互に配置することで、Jambaはメモリ使用量を削減しながら、高いパフォーマンスを維持します。

3. Mixture-of-Experts(MoE)モジュール

JambaのMoEモジュールは、モデルの容量をスケールアップするための柔軟なアプローチを提供します。MoEにより、モデルのパラメータ数を増やさずに、有効なパラメータ数を増やすことができます。Jambaでは、MoEは一部のMLP層に適用され、ルーター機構が各トークンにアクティブ化するトップエキスパートを選択します。この選択的なアクティブ化により、Jambaは高い効率性を維持しながら複雑なタスクを処理することができます。

以下の画像は、ハイブリッドAttention-Mambaモデルにおける誘導ヘッドの機能を示しています。ここでは、Attentionヘッドは、感情分析タスクに対する「Positive」または「Negative」などのラベルを予測する役割を果たします。強調表示された単語は、モデルのAttentionが、少数のショット例からのラベルトークンに強く焦点を当てていることを示しています。特に、最終的なラベルを予測する前に、モデルのAttentionは重要な役割を果たします。

MoEを統合することで得られるパフォーマンスの向上は、以下の表に示されています。MoEを使用することで、Jambaは計算コストを増やさずに容量を増やすことができます。これは、HellaSwag、WinoGrande、Natural Questions(NQ)などのベンチマークで大幅なパフォーマンスの向上に示されています。MoEを使用するモデルは、より高い精度(例:WinoGrandeで62.5%から66.0%)と、さまざまなドメインでのログ確率の向上(例:C4で-0.534)を達成します。

重要なアーキテクチャの特徴

  • 層の構成: Jambaのアーキテクチャは、Mamba層とTransformer層を特定の比率(例:1:7、つまり7つのMamba層ごとに1つのTransformer層)で組み合わせたブロックで構成されています。この比率は、パフォーマンスと効率の最適化のために調整されています。
  • MoEの統合: MoE層は、一定間隔で適用され、16人のエキスパートが利用可能で、各トークンごとにトップ2人のエキスパートがアクティブ化されます。この構成により、Jambaはメモリ使用量と計算効率のトレードオフを管理しながら、効果的にスケールアップすることができます。
  • 正規化と安定性: 訓練中の安定性を確保するために、JambaはMamba層にRMSNormを組み込みます。これにより、大規模なスケールでの大きなアクティベーションスパイクなどの問題を軽減することができます。

Jambaのパフォーマンスとベンチマーク

Jambaは、幅広いベンチマークで競争力のあるパフォーマンスを発揮しています。以下のセクションでは、Jambaが優れたパフォーマンスを示した主なベンチマークについて説明します。

1. 一般的なNLPベンチマーク

Jambaは、以下の学術的なベンチマークで評価されています。

  • HellaSwag(10ショット):Jambaは87.1%のパフォーマンススコアを達成し、多くの競合モデルを上回りました。
  • WinoGrande(5ショット):Jambaは82.5%のスコアを達成し、複雑な言語的推論を処理する能力を示しました。
  • ARC-Challenge(25ショット):Jambaは64.4%のスコアを達成し、難しい多肢選択問題を処理する能力を示しました。

MMLU(5ショット)などの集約ベンチマークでは、Jambaは67.4%のスコアを達成し、さまざまなタスクに対する堅牢性を示しました。

2. 長いコンテキストの評価

Jambaの特徴の1つは、非常に長いコンテキストを処理する能力です。モデルは、最大256Kトークンのコンテキスト長をサポートし、公開されているモデルの中で最も長いコンテキスト長をサポートしています。この機能は、Needle-in-a-Haystackベンチマークでテストされ、さまざまなコンテキスト長で優れたリトリーバル精度を示しました。

3. スループットと効率

Jambaのハイブリッドアーキテクチャは、特に長いシーケンスで、スループットを大幅に改善します。

さまざまなモデル間のスループット(トークン/秒)を比較するテストでは、Jambaは、特に大きなバッチサイズと長いコンテキストを使用するシナリオで、競合他社を一貫して上回りました。例えば、128Kトークンのコンテキストでは、JambaはMixtralと比較して3倍のスループットを達成しました。

Jambaの使用:Python

開発者や研究者がJambaを実験したい場合は、AI21 LabsはHugging Faceなどのプラットフォームでモデルを提供しています。以下のコードスニペットは、Jambaを使用してテキストを生成する方法を示しています。


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

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

<p>input_ids = tokenizer("In the recent Super Bowl LVIII,", return_tensors='pt').to(model.device)["input_ids"]</p>

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

print(tokenizer.batch_decode(outputs))

このシンプルなスクリプトは、Jambaモデルとトークナイザーをロードし、与えられた入力プロンプトに基づいてテキストを生成し、生成された出力を印刷します。

Jambaのファインチューニング

Jambaは、ベースモデルとして設計されており、特定のタスクやアプリケーションにファインチューニングできます。ファインチューニングにより、ユーザーはモデルをニッチなドメインに適応させ、専門タスクのパフォーマンスを向上させることができます。以下の例は、PEFTライブラリを使用してJambaをファインチューニングする方法を示しています。

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("ai21labs/Jamba-v0.1")
model = AutoModelForCausalLM.from_pretrained(
"ai21labs/Jamba-v0.1", device_map='auto', torch_dtype=torch.bfloat16)</p>

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

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

このコードスニペットは、英語の引用のデータセットでJambaをファインチューニングし、モデルを特定のタスクに適応させる方法を示しています。

デプロイと統合

AI21 Labsは、Jambaファミリーをさまざまなプラットフォームやデプロイオプションで提供しています。

  1. クラウドプラットフォーム
    • Google Cloud Vertex AI、Microsoft Azure、NVIDIA NIM (NVDA ) などの主要なクラウドプロバイダーで利用可能です。
    • Amazon Bedrock、Databricks Marketplace、Snowflake Cortexで近日公開予定です。
  2. AI開発フレームワーク:
    • LangChainやLlamaIndex(近日公開予定)などの人気フレームワークと統合されています。
  3. AI21 Studio:
    • AI21の独自開発プラットフォームから直接アクセスできます。
  4. Hugging Face
    • モデルはダウンロードして実験することができます。
  5. オンプレミスデプロイ:
    • 特定のセキュリティまたはコンプライアンス要件を持つ組織向けに、プライベートなオンプレミスデプロイオプションが提供されています。
  6. カスタムソリューション:
    • AI21は、エンタープライズクライアント向けにカスタマイズされたモデル調整とファインチューニングサービスを提供しています。

開発者向けの機能

開発者にとって、Jambaモデルは以下の機能を備えています。

  1. 関数呼び出し:外部ツールやAPIを簡単にAIワークフローに統合できます。
  2. 構造化JSON出力:自然言語入力から直接、クリーンで解析可能なデータ構造を生成できます。
  3. ドキュメントオブジェクト消化:複雑なドキュメント構造を効率的に処理して理解できます。
  4. RAG最適化:リトリーバル増強生成パイプラインを強化するための組み込み機能です。

これらの機能と、モデルの長いコンテキストウィンドウと効率的な処理により、Jambaは幅広い開発シナリオで多用途のツールとなります。

倫理的配慮と責任あるAI

while Jambaの機能は印象的ですが、その使用に責任あるAIの姿勢でアプローチすることが重要です。AI21 Labsは、以下の重要な点を強調しています。

  1. ベースモデルの性質:Jamba 1.5モデルは、特定のアライメントやインストラクションチューニングなしで事前トレーニングされたベースモデルです。
  2. 組み込みのセーフガードの欠如:モデルには、組み込みのモデレーションメカニズムがありません。
  3. 慎重なデプロイ:プロダクション環境またはエンドユーザーで使用する前に、追加の適応とセーフガードを実装する必要があります。
  4. データプライバシー:クラウドベースのデプロイを使用する場合は、データの取り扱いとコンプライアンス要件に注意する必要があります。
  5. バイアス認識:すべての大規模言語モデルと同様に、Jambaはトレーニングデータに存在するバイアスを反映する可能性があります。ユーザーはこれに気を付けて、適切な緩和策を実装する必要があります。

これらの要素を考慮することで、開発者や組織は、Jambaの機能を責任を持って倫理的に利用できます。

AI開発の新たな章

AI21 LabsがJambaファミリーを導入したことは、大規模言語モデルの進化における重要な里程標です。Transformerと状態空間モデルを組み合わせ、Mixture-of-Experts技術を統合し、コンテキスト長と処理速度の限界を押し進めることで、JambaはAIアプリケーションに新たな可能性を提供します。

AIコミュニティがこの革新的なアーキテクチャをさらに探求し、構築し続けるにつれて、モデル効率、長いコンテキストの理解、実用的AIデプロイのさらなる進歩が期待されます。Jambaファミリーは、単に新しいモデルセットではなく、大規模AIシステムの設計と実装の方法論の変化を表しています。

私は過去5年間、機械学習とディープラーニングの魅力的世界に没頭してきました。私の情熱と専門知識は、AI/MLに特に焦点を当てた50以上の多様なソフトウェアエンジニアリングプロジェクトに貢献することになりました。私の継続的な好奇心は、自然言語処理という分野にも私を引き付け、さらに探求したいと思っています。