pile·
AI / ML·스캐터랩스캐터랩 (이루다)·

더 나은 생성모델을 위해 RLHF로 피드백 학습시키기

대형 생성 모델을 사람의 의도에 맞게 정렬하는 RLHF의 원리(SFT→리워드 모델→PPO)와, 이를 루다에 실제 적용한 결과·한계·최적화 기법을 다룬다. 루다는 별도 리워드 모델 대신 기존 랭킹 모델을 리워드로 재활용했다.

핵심 포인트
  • RLHF는 SFT로 답변 형태를 잡고, 사람이 매긴 선호 순위로 리워드 모델을 학습한 뒤, PPO로 리워드를 최대화하도록 파인튜닝한다.
  • KL penalty로 참조 모델 분포에서 너무 벗어나지 않게 막아 리워드 해킹·mode collapse를 억제한다.
  • 루다는 별도 리워드 모델 없이 어뷰징·선정성·적합성·재미를 평가하던 기존 랭킹 모델 점수를 리워드로 쓰고, trlX를 참고해 PPO 학습했다.
  • 결과 평균 리워드가 1.23→1.36으로 올랐지만, 2500 step에선 괄호문·부자연스러운 답변 등 mode collapse가 나타났다.
  • 3B 모델 RLHF는 4개 모델로 총 96GiB 메모리가 필요해 FSDP·LoRA·Flash Attention 등 최적화가 필수다.
상세 정리
  • 학습 3단계: Pre-training(의도대로 동작 어려움) → SFT(정제·레이블링 데이터로 원하는 답변 형태 학습) → RLHF(사람 피드백 기반 강화학습).
  • 리워드 모델: 문맥에 SFT가 만든 후보들을 사람이 선호 순위로 레이블링하고, Bradley-Terry로 positive logit은 키우고 negative는 낮추도록 학습한다.
  • RLHF 학습: 별도 문맥에서 후보 생성 후 리워드를 계산하고 PPO로 리워드를 최대화한다. KL penalty는 참조 분포 이탈을 막는 regularization이다.
  • 루다 적용: 랭킹 모델이 최적 답변을 고르는 구조였는데, RLHF로 생성 모델이 스스로 좋은 답변을 내면 서빙 때 랭킹 모델을 빼 구조가 단순해진다. 기존 랭킹 점수를 리워드로 재사용했다.
  • 결과: 왼쪽 그래프는 리워드 상승(평균 1.23→1.36), KL 그래프는 일부 step을 빼면 분포 이탈이 크지 않았다. 다만 1500 step은 자연스럽지만 2500 step은 지시문형 괄호문 등 부자연스러움(리워드 해킹)이 나타났다.
  • 메모리 최적화: 3B 기준 4모델 가중치 24GiB(BFloat16 2Byte/param) + optimizer·gradient 72GiB = 96GiB(activation 제외). PyTorch FSDP로 가중치를 sharding해 GPU 간 공유하고, 낮은 정밀도·LoRA(PEFT)·파라미터 공유를 병행한다.
  • 추론 최적화: 학습 중 문장 생성이 있어 Flash Attention·key-value caching으로 샘플링을 가속한다.
  • 불안정성 대응: 리워드 모델 강건성이 부족하면 mode collapse가 난다. 리워드 모델 확대, negative sample augmentation, pretraining-mix, EMA 적용으로 완화한다.
  • 대안: 더 적은 모델로 안정적인 Rejection Sampling(Best-of-N), RRHF(랭킹 loss, 사람·GPT-4·ChatGPT·SFT 샘플), DPO(리워드 모델 없이 참조 모델로 선호쌍 직접 학습, RLHF보다 안정적).
  • 결론: RLHF로 루다의 올바른 답변 생성 개선 가능성을 확인했다. 강건한 리워드 모델이 있으면 추가 레이블 없이 continual learning이, 나아가 사람 없이 AI 피드백을 쓰는 RLAIF도 연구 대상이다.
왜 읽나자사 생성 모델에 RLHF를 얹으려는 ML 엔지니어에게 랭킹 모델 재활용·PPO·메모리 최적화(FSDP/LoRA)·mode collapse 대응 실전 레퍼런스.
스캐터랩
스캐터랩 (이루다) 블로그
원문은 여기서 이어서 읽을 수 있어요
원문 읽기
읽음 (0)

이 글과 비슷한

  1. AI / ML·LY CorporationLY Corporation·

    Grafana에서 자연어로 장애 원인을 분석하기: LLM 에이전트 기반 SRELens 개발기

    LY Corporation Home SRE 팀이 장애 분석 시 메트릭·로그·트레이스가 각각 다른 화면에 흩어져 있는 문제를 해결하기 위해 Grafana 플러그인 SRELens를 개발했다. SRELens는 LLM 에이전트가 자연어 질의를 받아 실제 관측성 데이터를 조회하고, 근거와 함께 장애 원인 후보를 정리해 주는 도구다. LGTM-P 스택(Loki·Grafana·Tempo·Mimir·Pyroscope)과 FlavaMCP 게이트웨이를 통합해 단일 채팅 인터페이스에서 멀티시그널 분석이 가능하다.

    #llm-app#mcp#observability+2