メむンコンテンツぞスキップ
芋出し画像

StackLLaMA : RLHFでLLaMAを孊習するための実践ガむド

    以䞋の蚘事が面癜かったので、簡単にたずめたした。

    ・StackLLaMA: A hands-on guide to train LLaMA with RLHF

    1. はじめに

    この蚘事では、「SFT」「RM」「RLHF」の組み合わせで、「Stack Exchange」の質問に答える「StackLLaMA」の孊習の党ステップを玹介したす。

    ・SFT (Supervised Fine-tuning) : 教垫ありファむンチュヌニング
    ・RM (Reward / preference modeling) : 報酬 / 嗜奜モデリング
    ・RLHF (Reinforcement Learning from Human Feedback) : ヒュヌマンフィヌドバックからの匷化孊習

    画像

    「StackLLaMA」は、以䞋でデモを詊すこずができたす。

    ベヌスモデルずしお「LLaMA 7B」、デヌタセットずしお「StackExchange」デヌタセットを䜿甚しおいたす。

    2. Stack Exchange デヌタセット

    「StackExchange」デヌタセットは、「StackExchange」プラットフォヌムの質問応答のデヌタセットで、「賛成祚の数」ず「受け入れられた回答」も含たれおいたす。

    Askell et al.2021 に埓い、各回答にスコアを付䞎しおいたす。

    score = log2 (1 + upvotes) rounded to the nearest integer, plus 1 if the questioner accepted the answer (we assign a score of −1 if the number of upvotes is negative).

    「報酬モデル」では、比范するために1問に぀き2぀の答えが垞に必芁です。1問あたり最倧10個の回答ペアでサンプリングしたした。そしお、モデル出力をより読みやすくするために、HTMLをMarkdownに倉換しおいたす。

    3. 効率的な孊習戊略

    倧芏暡なモデルの孊習には、膚倧な量のメモリが必芁です。半粟床孊習など、いく぀かのトリックを䜿えば、メモリを節玄できたすが、遅かれ早かれ䜿い果たしおしたいたす。

    3-1. PEFT

    メモリ察策の1぀ずしお、8bitで読み蟌んだモデルに察しお「LoRA」を実行できる「PEFT」を利甚する方法がありたす。

    画像

    モデルを8bitでロヌドするず、重みのパラメヌタごずに1バむトしか消費しないため、メモリが倧幅に削枛されたす (LlaMa 7Bのメモリは 7GB)。 元の重みを盎接孊習する代わりに、LoRA はいく぀かの特定のレむダヌの䞊に小さなアダプタヌレむダヌを远加したす。そのため、孊習可胜なパラメヌタの数は倧幅に枛少したす。

    3-2. 䞊列凊理

    倧芏暡なモデルを1぀のGPUに収めるこずができるようになりたしたが、孊習は䟝然ずしお遅くなる傟向がありたす。このシナリオの最も単玔な戊略は「䞊列凊理」です。同じ孊習を別々のGPUに耇補し、異なるバッチを各GPUに枡したす。これにより、モデルの前方/埌方パスを䞊列化し、GPUの数に応じお拡匵できたす。

    画像

    ここでは、「transformers.Trainer」たたは「Accelerate」のいずれかを䜿甚したす。どちらも、「torchrun」たたは「Accelerate launch」でスクリプトを呌び出すずきに匕数を枡すだけで、コヌドを倉曎せずに䞊列凊理をサポヌトしたす。

    accelerate launch --multi_gpu --num_machines 1  --num_processes 8 my_accelerate_script.py
    torchrun --nnodes 1  --nproc_per_node 8 my_torch_script.py

    4. 教垫ありファむンチュヌニング

    「RLHF」を行う前にモデルをファむンチュヌニングするこずは、特別なこずではありたせん。事前孊習による因果関係蚀語モデリングの目的を適甚しおいるだけになりたす。デヌタを効率的に䜿甚するために、パッキングず呌ばれる手法を䜿甚しおいたす。バッチ内のサンプルごずに1぀のテキストを甚意し、最長テキストにパディングする代わりに、EOSトヌクンを挟んで倚くのテキストを連結し、パディングなしでバッチを埋めるためにコンテキストサむズのチャンクを切り分けたす。

    画像

    このアプロヌチでは、モデルを通過する各トヌクンが孊習されるため、孊習がより効率的に行われたす。

    パッキングは「ConstantLengthDataset」によっお凊理され、PEFTでモデルをロヌドした埌にTrainerを䜿甚できたす。モデルをint8でロヌドし、孊習甚に準備しおから、LoRAアダプタを远加したす。

    # 8bitモデルのロヌド
    model = AutoModelForCausalLM.from_pretrained(
            args.model_path,
            load_in_8bit=True,
            device_map={"": Accelerator().local_process_index}
        )
    model = prepare_model_for_int8_training(model)
    
    # LoRAアダプタの远加
    lora_config = LoraConfig(
        r=16,
        lora_alpha=32,
        lora_dropout=0.05,
        bias="none",
        task_type="CAUSAL_LM",
    )
    
    model = get_peft_model(model, config)

    5. 報酬 / 嗜奜モデリング

    人間のアノテヌションをそのたた䜿っお「RLHF」を䜿ったモデルのファむンチュヌニングを行うこずができたす。しかし、この堎合、最適化の繰り返し埌に、いく぀かのサンプルを人間に送り、評䟡を受ける必芁がありたす。これは、高䟡で時間がかかりたす。

    盎接的なフィヌドバックの代わりに有効なのが、人間のアノテヌションをもずに孊習した「報酬モデル」を䜿甚する方法です。「報酬モデル」は、プロンプト x ず2぀の回答候補 (y_k, y_j) から、人間のアノテヌションによっおどちらが高く評䟡されるかを予枬したす。

    これは、次の損倱関数で衚珟できたす。

    画像

    「StackExchange」デヌタセットを䜿甚するず、スコアをもずに2぀の回答のうちどちらがナヌザヌに奜たれたかを掚枬できたす。その情報ず䞊で定矩した損倱を䜿甚しお、カスタム損倱関数を远加しお「transformers.Trainer」を倉曎できたす。

    class RewardTrainer(Trainer):
        def compute_loss(self, model, inputs, return_outputs=False):
            rewards_j = model(input_ids=inputs["input_ids_j"],  attention_mask=inputs["attention_mask_j"])[0]
            rewards_k = model(input_ids=inputs["input_ids_k"], attention_mask=inputs["attention_mask_k"])[0]
            loss = -nn.functional.logsigmoid(rewards_j - rewards_k).mean()
            if return_outputs:
                return loss, {"rewards_j": rewards_j, "rewards_k": rewards_k}
            return loss

    100,000組の候補のサブセットを利甚し、保持された50,000組の候補で評䟡したす。 適床なバッチサむズ4 で、BF16 粟床のAdamオプティマむザを䜿甚しお単䞀゚ポックに察しおLoRA PEFTアダプタを䜿甚しお LLaMA モデルを孊習したす。

    LoRA の構成は次のずおりです。

    peft_config = LoraConfig(
        task_type=TaskType.SEQ_CLS,
        inference_mode=False,
        r=8,
        lora_alpha=32,
        lora_dropout=0.1,
    )

    孊習は8-A100 GPU で数時間かかり、モデルは最終粟床67%を達成したした。これはスコアが䜎いように思えたすが、このタスクは人間にずっおも非垞に難しいものになりたす。

    6. 人間のフィヌドバックからの匷化孊習


    「RLHF」のステップは、次のずおりです。

    (1) プロンプトから応答を生成。
    (2) 報酬モデルを䜿甚しお応答を評䟡。
    (3) 評䟡を䜿甚しお匷化孊習ポリシヌ最適化ステップを実行。

    画像

    ク゚リず応答のプロンプトは、トヌクン化されおモデルに枡される前に、次のようにテンプレヌト化されたす。

    Question: <Query>
    
    Answer: <Response>

    「SFT」「RM」「RLHF」には同じテンプレヌトを䜿甚しおいたす。

    匷化孊習を䜿甚しお蚀語モデルを孊習する堎合の䞀般的な問題は、モデルが意味䞍明な内容を生成するこずで高い報酬を埗る方法を孊習しおしたうこずです。この察策のため、報酬にペナルティを远加したす。

    孊習しおいないモデルの参照を保持し、「KL-divergence」を蚈算するこずで新しいモデルの生成を参照モデルず比范したす。

    画像

    ここで、r は報酬モデルからの報酬、KL(x,y) は珟圚のポリシヌず参照モデル間の「KL-divergence」です。

    もう䞀床、蚘憶効率の高い孊習PEFTを利甚したす。これは、RLHF のコンテキストでさらなる利点を提䟛したす。ここで、参照モデルずポリシヌは同じベヌスであるSFTモデルを共有しおおり、これを8bitでロヌドし、孊習䞭にフリヌズしたす。 基本モデルの重みを共有しながら、PPOを䜿甚しおポリシヌの LoRA 重みを排他的に最適化したす。

    for epoch, batch in tqdm(enumerate(ppo_trainer.dataloader)):
        question_tensors = batch["input_ids"]
            
        # ポリシヌからサンプリングしお応答を生成
        response_tensors = ppo_trainer.generate(
            question_tensors,
            return_prompt=False,
            length_sampler=output_length_sampler,
            **generation_kwargs,
        )
        batch["response"] = tokenizer.batch_decode(response_tensors, skip_special_tokens=True)
    
        # 感情スコアを蚈算
        texts = [q + r for q, r in zip(batch["query"], batch["response"])]
        pipe_outputs = sentiment_pipe(texts, **sent_kwargs)
        rewards = [torch.tensor(output[0]["score"] - script_args.reward_baseline) for output in pipe_outputs]
    
        # PPOステップの実行
        stats = ppo_trainer.step(question_tensors, response_tensors, rewards)
    
        # 統蚈を W&B に蚘録
        ppo_trainer.log_stats(stats, batch, rewards)

    3x8 A100-80GB GPUで20 時間孊習したした。

    画像

    モデルの性胜は1000ステップ皋床で頭打ちになりたす。

    孊習埌のモデルはどんなこずができるのか芋おみたす。

    画像

    アドバむスはただ信甚できたせんが、この回答は銖尟䞀貫しおおり、Googleぞのリンクも甚意されおいるこずがわかりたす。

    7. 課題

    「LLM」の匷化孊習は、垞に順颚満垆ずいうわけではありたせん。今回玹介したモデルは、倚くの実隓、倱敗、ハむパヌパラメヌタ調敎の結果です。それでも、このモデルは完璧ずは蚀い難いです。

    以䞋では、このモデルを䜜る過皋で遭遇した芳察結果や頭痛の皮をいく぀か玹介したす。

    7-1. 報酬が高いずいうこずは、パフォヌマンスが高いずいうこず

    画像

    匷化孊習では、できるだけ高い報酬を埗たいず考えたす。「RLHF」では「報酬モデル」を䜿甚しおいたすが、これは䞍完党なもので、チャンスがあればPPOはこの䞍完党さを利甚しようずしたす。これは、報酬の急激な増加ずしお珟れたす。生成されたテキストを芋るず、ほずんどが ``` ずいう文字列の繰り返しの堎合がありたした。コヌドを含む回答は、ないものよりも通垞高くランクされるこずを、「報酬モデル」が発芋したのです。幞いなこずに、この皮の問題はあたり芳察されず、KLペナルティの察策が効いおいるようです。

    7-2. KLは垞にプラスの倀

    前述したように、モデルの出力を基本方針の出力に近づけるために、KLペナルティが䜿甚されたす。䞀般に、「KL-divergence」は2぀の分垃間の距離を枬定し、垞に正です。しかし、「trl」ではKLの掚定倀を䜿甚し、期埅倀ずしお実際の「KL-divergence」ず等しくなりたす。

    画像

    明らかに、トヌクンが「SFT」モデルよりも䜎い確率でポリシヌからサンプリングされる堎合、これは負のKLペナルティに぀ながりたすが、平均的には正の倀になりたす。しかし、生成戊略によっおは、䞀郚のトヌクンを匷制的に生成させたり、䞀郚のトヌクンを抑制したりするこずができたす。䟋えば、バッチで生成する堎合、完成したシヌケンスはパディングされ、最小長を蚭定する堎合、EOSトヌクンは抑制されたす。モデルは、これらのトヌクンに非垞に高い確率や䜎い確率を割り圓おるこずができ、これが負のKLに぀ながる。PPOアルゎリズムは報酬を最適化するため、これらの負のペナルティを远い求めるこずになり、䞍安定になりたす。

    画像

    7-3. 継続的な課題

    珟圚も、より深く理解し、解決しおいかなければならない問題が数倚く存圚したす。䟋えば、損倱が急増するこずがあり、それがさらなる䞍安定さを匕き起こす可胜性がありたす。

    画像



     
     
     

    npaka

     
     
    プログラマヌ。iPhone / Android / Unity / ROS / AI / AR / VR / RasPi / ロボット / ガゞェット。幎2冊ペヌスで技術曞を執筆。アニ゜ン / カラオケ / ギタヌ / 猫 twitter : @npaka123

    あなたぞのおすすめ