Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 34 additions & 0 deletions .github/workflows/trl-quick-start.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,34 @@
# trl quick start guard - thin trigger over quick-start-template.yml; the engine owns the monitor loop, cache I/O and result publishing.
name: trl-quick-start

concurrency:
group: ${{ github.event_name == 'schedule' && 'trl-quick-start-schedule' || format('manual-{0}', github.run_id) }}
cancel-in-progress: false

on:
schedule:
- cron: '0 */3 * * *'
workflow_dispatch:
pull_request:
branches: [main]
paths:
- 'sources/trl/**'
- 'tests/trl/**'

permissions:
contents: read

jobs:
trl-quick-start:
uses: ./.github/workflows/quick-start-template.yml
with:
project: trl
test_runner: '["linux-aarch64-a2-1"]'
image: swr.cn-south-1.myhuaweicloud.com/ascendhub/cann:9.1.0-910b-ubuntu22.04-py3.12
container_options: >-
--volume=/data/ci-cache/modelscope/trl:/root/.cache/modelscope
timeout_minutes: 180
upstream_repo: huggingface/trl
doc_url: 'https://raw.githubusercontent.com/Ascend/docs/{0}/sources/trl/quick_start.md'
doc_path: https://github.com/Ascend/docs/blob/main/sources/trl/quick_start.md
test_command: python -m unittest tests.trl.test_quick_start_ascend -v 2>&1
3 changes: 2 additions & 1 deletion conf.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,7 +77,8 @@
'sources/llama_cpp/quick_start.md',
'sources/whisper_cpp/quick_start.md',
'sources/llm_compressor/quick_start.md',
'sources/axolotl/quick_start.md']
'sources/axolotl/quick_start.md',
'sources/trl/quick_start.md']


# -- Options for HTML output -------------------------------------------------
Expand Down
2 changes: 1 addition & 1 deletion index.rst
Original file line number Diff line number Diff line change
Expand Up @@ -214,7 +214,7 @@
<div class="project-card">
<div class="card-top"><div class="card-icon" style="background-image: url('_static/images/huggingface.png')"></div><h3 class="card-title">Transformer Reinforcement Learning</h3></div>
<p class="card-desc">适用于 SFT、PPO、DPO 等方法的模型后训练库。</p>
<div class="card-footer"><a href="https://github.com/huggingface/trl">官方链接</a><span class="split">|</span><a href="sources/trl/install.html">安装指南</a><span class="split">|</span><a href="sources/trl/quick_start.html">快速上手</a></div>
<div class="card-footer"><a href="https://github.com/huggingface/trl">官方链接</a><span class="split">|</span><a href="sources/trl/index.html">快速上手</a></div>
</div>

<!-- Twinkle:官方文档站已含 NPU 说明,外链跳转,不再本地编译 -->
Expand Down
Binary file removed sources/trl/images/image.png
Binary file not shown.
10 changes: 2 additions & 8 deletions sources/trl/index.rst
Original file line number Diff line number Diff line change
@@ -1,8 +1,2 @@
Transformer Reinforcement Learning
===================================================

.. toctree::
:maxdepth: 2

install.rst
quick_start.rst
.. include:: quick_start.md
:parser: myst_parser.sphinx_
38 changes: 0 additions & 38 deletions sources/trl/install.rst

This file was deleted.

191 changes: 191 additions & 0 deletions sources/trl/quick_start.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,191 @@
# TRL

TRL 用统一的 `Trainer` / `Config` API 支持多种模型后训练方法。本示例在单卡昇腾 NPU 上,用 Qwen2.5-0.5B-Instruct 分别运行 SFT 和 DPO LoRA。

## 前置条件

### 硬件

Atlas 900 A2 / A3 训练系列产品或者 Ascend 950 系列产品,并按需完成物理机或容器内的设备挂载。

### 基础软件

在运行本文档示例之前,你的机器上需要已经装好并可用:

- 可用的 Python 环境
- 可用的 CANN(参考[快速安装昇腾环境](https://ascend.github.io/docs/sources/ascend/quick_install.html))
- 根据 CANN 版本安装匹配的 `torch_npu`(参考 [Ascend PyTorch 安装文档](https://gitcode.com/Ascend/pytorch))

本文档示例在 Python 3.12、CANN 9.1.0、`torch_npu` 2.9.0.post2 环境下验证通过。

## 加载 CANN 环境

```shell
source /usr/local/Ascend/ascend-toolkit/set_env.sh
```

## 安装 TRL

安装 TRL 并查看安装版本:

```shell #test id="install-trl"
python -m pip install trl
python -c "import trl; print('trl', trl.__version__)"
```

输出结果如下,其中 `xxx` 为实际安装的 TRL 版本:

```shell #test-result id="install-trl" fuzzy='...' fuzzy='xxx'
...
trl xxx
```

## 示例一:SFT LoRA 后训练

用 Qwen2.5-0.5B-Instruct 和 ModelScope 的 `HuggingFaceH4/ultrafeedback_binarized` SFT 子集进行 5 步 LoRA 微调,模型与数据集会自动下载,适配器保存到 `output/trl-sft-lora`。

安装示例依赖:

```shell #test-setup
python -m pip install peft "transformers>=4.56.2,<5.0" datasets "modelscope==1.37.0"
```
运行 SFT 训练脚本:
```python #test id="sft-lora"
import os
import shutil
import torch
import torch_npu
from datasets import load_dataset
from modelscope import snapshot_download
from peft import LoraConfig, TaskType
from trl import SFTConfig, SFTTrainer

print("TRL_SFT_BEGIN")

ds_path = snapshot_download(
'HuggingFaceH4/ultrafeedback_binarized', repo_type='dataset',
)
data_dir = './ultrafeedback_sft'
if os.path.isdir(data_dir):
shutil.rmtree(data_dir)
os.makedirs(data_dir, exist_ok=True)
for name in os.listdir(os.path.join(ds_path, 'data')):
if name.startswith('train_sft-'):
shutil.copy2(os.path.join(ds_path, 'data', name), data_dir)
train_dataset = load_dataset(
'parquet', data_files=os.path.join(data_dir, 'train_sft-*.parquet'),
split='train',
).select_columns(['messages'])

model = snapshot_download('Qwen/Qwen2.5-0.5B-Instruct')

trainer = SFTTrainer(
model=model,
train_dataset=train_dataset,
peft_config=LoraConfig(r=8, lora_alpha=32, task_type=TaskType.CAUSAL_LM),
args=SFTConfig(
output_dir="output/trl-sft-lora",
max_steps=5,
per_device_train_batch_size=1,
gradient_accumulation_steps=1,
learning_rate=1e-4,
max_length=512,
logging_steps=1,
save_strategy="no",
report_to="none",
model_init_kwargs={"dtype": torch.bfloat16},
),
)
print("model device:", next(trainer.model.parameters()).device)
trainer.train()
trainer.save_model("output/trl-sft-lora")
print("LoRA adapter saved to: output/trl-sft-lora")
print("TRL_SFT_DONE")
```

输出结果类似如下:

```shell #test-result id="sft-lora"
...
LoRA adapter saved to: output/trl-sft-lora
TRL_SFT_DONE
```

## 示例二:偏好优化 DPO LoRA

再用相同模型和数据集运行 3 步 DPO LoRA,适配器保存到 `output/trl-dpo-lora`。

```python #test id="dpo-lora"
import os
import shutil
import torch
import torch_npu
from datasets import load_dataset
from modelscope import snapshot_download
from peft import LoraConfig, TaskType
from transformers import AutoModelForCausalLM, AutoTokenizer
from trl import DPOConfig, DPOTrainer

print("TRL_DPO_BEGIN")

ds_path = snapshot_download(
'HuggingFaceH4/ultrafeedback_binarized', repo_type='dataset',
)
data_dir = './ultrafeedback_prefs'
if os.path.isdir(data_dir):
shutil.rmtree(data_dir)
os.makedirs(data_dir, exist_ok=True)
for name in os.listdir(os.path.join(ds_path, 'data')):
if name.startswith('train_prefs-'):
shutil.copy2(os.path.join(ds_path, 'data', name), data_dir)
train_dataset = load_dataset(
'parquet', data_files=os.path.join(data_dir, 'train_prefs-*.parquet'),
split='train',
)

# prompt 列是纯字符串,chosen / rejected 是 messages 列表;把 prompt 转成
# 单条 user 消息即可让 DPOTrainer 按 conversational 格式处理
def to_conversational(example):
example['prompt'] = [{'role': 'user', 'content': example['prompt']}]
return example

train_dataset = train_dataset.map(to_conversational)

model_path = snapshot_download('Qwen/Qwen2.5-0.5B-Instruct')
model = AutoModelForCausalLM.from_pretrained(model_path, dtype=torch.bfloat16)
tokenizer = AutoTokenizer.from_pretrained(model_path)

trainer = DPOTrainer(
model=model,
ref_model=None,
processing_class=tokenizer,
train_dataset=train_dataset,
peft_config=LoraConfig(r=8, lora_alpha=32, task_type=TaskType.CAUSAL_LM),
args=DPOConfig(
output_dir="output/trl-dpo-lora",
max_steps=3,
per_device_train_batch_size=1,
gradient_accumulation_steps=1,
learning_rate=1e-4,
max_length=512,
logging_steps=1,
save_strategy="no",
report_to="none",
),
)
print("model device:", next(trainer.model.parameters()).device)
trainer.train()
trainer.save_model("output/trl-dpo-lora")
print("LoRA adapter saved to: output/trl-dpo-lora")
print("TRL_DPO_DONE")
```

输出结果类似如下:

```shell #test-result id="dpo-lora"
...
LoRA adapter saved to: output/trl-dpo-lora
TRL_DPO_DONE
```

更多方法(GRPO / PPO / Reward / KTO 等)见 [TRL examples](https://github.com/huggingface/trl/tree/main/examples)。
49 changes: 0 additions & 49 deletions sources/trl/quick_start.rst

This file was deleted.

13 changes: 13 additions & 0 deletions tests/trl/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,13 @@
"""Tests package marker (injects repo tests/ into sys.path)."""

from __future__ import annotations

import sys
from pathlib import Path

_REPO_ROOT = Path(__file__).resolve().parents[2]
_TESTS_ROOT = _REPO_ROOT / 'tests'
for _p in (_TESTS_ROOT, _REPO_ROOT):
_ps = str(_p)
if _ps not in sys.path:
sys.path.insert(0, _ps)
Loading
Loading