Pass Rate 필터링 — 데이터셋을 검증 가능한 문제로 거르기

Pass Rate 필터링 — 데이터셋을 검증 가능한 문제로 거르기

합성 데이터로 학습할 때는 품질이 좋은 문제만 남기는 게 중요해요. Open R1은 검증 가능한 태스크에서 pass rate(합격률)를 계산해 데이터셋을 필터링하는 지원을 제공해요.

출처: https://github.com/huggingface/open-r1/blob/main/scripts/pass_rate_filtering/README.md

스크립트는 scripts/pass_rate_filtering/compute_pass_rate.pyscripts/pass_rate_filtering/launch_filtering.sh이며, 현재는 DAPO 하드코딩돼 있어요. 기본적으로 데이터셋을 청크 단위로 나누고, 합치려면 다음과 같이 실행해요 (DAPO 예시).

from datasets import load_dataset, concatenate_datasets

name = "open-r1/DAPO-Math-17k-Processed-R1-Distill-Qwen-Math-7B-Merges-v00.02-v01.02-0.3-0.7-filter"
gen_datasets = []
filt_datasets = []
for start in range(0,17400,200):
    end = start + 200
    if start == 17200:
        end = 17398
    gen_config_name = f"gen-{start}-{end}"
    gen_dataset = load_dataset(name, gen_config_name, revision="gen",  split="train")
    gen_datasets.append(gen_dataset)
    filt_config_name = f"filt-0.1-0.6-{start}-{end}"
    filt_dataset = load_dataset(name, filt_config_name, revision="pass_rate",  split="train")
    filt_datasets.append(filt_dataset)
gen_dataset = concatenate_datasets(gen_datasets)
gen_dataset.push_to_hub(name, config_name="gen", split="train")
filt_dataset = concatenate_datasets(filt_datasets)
filt_dataset.push_to_hub(name, config_name="default", split="train")
gen-{start}-{end}
filt-0.1-0.6-{start}-{end}

이렇게 생성(gen) 버전과 pass rate 필터(filt) 버전을 각각 병합해 Hub에 다시 올리면 정리된 학습 데이터셋을 얻을 수 있어요.

더 알아보기