Gemma4 MTP DraftersをHuggingFaceから試す

お疲れ様です。

先月Gemma4の発表に続き、Gemma4 MTP Draftersがリリースされました。 今回はこちらを試していきたいと思います。

先月のGemma4:E4Bを試したときの記事も併せてどうぞ。

fallpoke-tech.hatenadiary.jp

Gemma4 MTP Drafters とは

ざっくりと説明すると、Gemma4のモデルに加えて補助用のモデルを用いて高速化したものです。 具体的な仕組みについては以下Google公式のリリースや解説記事を参照ください。

blog.google

note.com

実装

ソースコードは例によってこちらの検証環境のGitHubリポジトリに残してあります。

github.com

今回のモデル(というか手法?)はOllamaで使うことができなかったので、HuggingFaceで公開されているモデルを利用します。

ソースコードは以下になります。リポジトリ内の"chat/generate_response_gemma4.py"です。
注目すべきはTarget ModelとAssistant Modelを個別に呼び出している部分です。 Target Modelで返答を生成するときにAssistant Modelを引数で設定しています。

2つのモデルを使用しているせいか、少しGPUメモリの使用量は増えているように思います。 前回まででGemma4:E4Bを使用していましたそちらでは私の実行環境のGPU(RTX4060ti 16GB)に乗り切らなかったのでそれより軽量のモデルであるGemma4:E2Bの方を使用しています。

# Target Model
processor = AutoProcessor.from_pretrained("google/gemma-4-E2B-it")
target_model = AutoModelForCausalLM.from_pretrained(
    "google/gemma-4-E2B-it",
    dtype="auto",
    device_map="auto",
)

# Assistant Model (the drafter)
assistant_model = AutoModelForCausalLM.from_pretrained(
    "google/gemma-4-E2B-it-assistant",
    dtype="auto",
    device_map="auto",
)

def response_generator_langchain_gemma4_rag() -> str:
    """langchain_huggingfaceを使用+RAG(Gemma4用)
    """    
    # 直前のユーザの入力を取得
    user_input = st.session_state.messages[-1]["content"]
    
    with open("./biography_context.txt", mode="r", encoding="utf-8") as f:
        context = f.read()
    
    # システムプロンプトの用意
    system_prompt = [{
        "role": "system",
        "content": "あなたはユーザの質問に答えるアシスタントです。"
    }]
    
    # 直前のユーザの入力を取得
    rag_input = [{
        "role": "user",
        "content": RAG_PROMPT.format(question=user_input, context=context)
    }]
    
    # Process input
    text = processor.apply_chat_template(
        system_prompt + st.session_state.messages[:-1] + rag_input, 
        tokenize=False, 
        add_generation_prompt=True, 
    )
    inputs = processor(text=text, return_tensors="pt").to(target_model.device)
    input_len = inputs["input_ids"].shape[-1]
    
    # Generate output
    outputs = target_model.generate(
        **inputs,
        assistant_model=assistant_model,
        max_new_tokens=1024,
    )
    response = processor.decode(outputs[0][input_len:], skip_special_tokens=False)
    
    # Parse output
    parsed_response = processor.parse_response(response)
    
    response_html = Markdown().convert(parsed_response["content"])
    
    return response_html

実行

上記のコードを使用して実行しました。
RAGのコードになっており、コンテキストはエレファントカシマシの公式サイトのBiographyの内容をすべて与えています。

出力結果は以下のようになりました。
モデルとしてはGemma4シリーズで一番小さいですがいい感じの出力になっていると思います。 やっぱりGemma4シリーズは日本語の性能が高いように感じますね。

output

注目の実行時間ですが、なんと5秒ほどでした。
モデルサイズや実行環境など条件が異なるのですが、前回Gemma4:E4Bを試した時は最大で45秒ほどかかっていたのでなかなかに衝撃でした。

ローカルLLMでもここまでの速さで精度の良い出力できるようになったのはすごいですね。

elapsed_time