DeepSeek用的GRPO占用大量内存?有人给出了些破解方法

选自oxen.ai

作者:Greg Schoeninger

编译:陈陈、泽南

RTX 3080 移动版能训练哪种大模型?本文为那些 GPU 资源有限时使用 GRPO 训练的开发者提供了宝贵的指导。

自 DeepSeek-R1 发布以来,群组相对策略优化(GRPO)因其有效性和易于训练而成为大型语言模型强化学习的热门话题。R1 论文展示了如何使用 GRPO 从遵循 LLM(DeepSeek-v3)的基本指令转变为推理模型(DeepSeek-R1)。

GRPO 是一种在线学习算法(online learning algorithm),它通过使用训练过程中由训练模型自身生成的数据来进行迭代改进。GRPO 的目标是最大化生成补全(completions)的优势函数(advantage),同时确保模型保持在参考策略(reference policy)附近。

本文的目的是帮你节省一些时间,让你根据硬件预算选择合适的模型大小。在开始微调时,你必须做出的重要决定是选择模型大小,以及你是执行完全微调还是参数高效微调(PEFT)。

文章作者来自 AI 公司 Oxen.ai 的 CEO Greg Schoeninger。

原文链接:https://www.oxen.ai/blog/grpo-vram-requirements-for-the-gpu-poor

作者表示,他发现 trl 库中已经有一个易于使用的 GRPO 实现,便立刻开始了训练,使用的硬件是配备了 16GB 显存的 Nvidia GeForce RTX 3080 的小型笔记本电脑。正如大家可能遇到的问题,作者发现示例代码中的参数设置导致了一个巨大的显存不足(OOM,out of memory )错误。

  1. torch

  2. .

  3. OutOfMemoryError

  4. :

  5. CUDA

  6. out

  7. of memory

  8. .

  9. Tried

  10. to allocate

  11. 1.90


  12. GiB

  13. .

  14. GPU

  15. 0

  16. has a total capacity of

  17. 15.73


  18. GiB

  19. of which

  20. 1.28


  21. GiB


  22. is

  23. free

  24. .


  25. Including

  26. non

  27. -

  28. PyTorch

  29. memory

  30. ,


  31. this

  32. process has

  33. 14.43


  34. GiB

  35. memory

  36. in


  37. use

  38. .


  39. Of

  40. the allocated memory

  41. 11.82


  42. GiB


  43. is

  44. allocated

  45. by


  46. PyTorch

  47. ,


  48. and


  49. 2.41


  50. GiB


  51. is

  52. reserved

  53. by


  54. PyTorch

  55. but unallocated

  56. .


  57. If

  58. reserved but unallocated memory

  59. is

  60. large

  61. try

  62. setting PYTORCH_CUDA_ALLOC_CONF

  63. =

  64. expandable_segments

  65. :

  66. True

  67. to avoid fragmentation

  68. .


  69. See

  70. documentation

  71. for


  72. Memory


  73. Management


  74. (

  75. https

  76. :

  77. //pytorch.org/docs/stable/notes/cuda.html#environment-variables)

实际使用情况

作者表示,他们进行了一系列实验,以确定训练各种大小的模型所需的显存(VRAM)要求。参数数量从 5 亿到 140 亿不等,他们比较了权重的完全微调与参数高效微调(使用 LoRA),所有训练运行都在英伟达 H100 上完成,因此这里的 OOM 意味着 >80GB 的 VRAM。

在表格中,你可以找到 GSM8K 数据集上训练的前 100 步中的峰值内存使用情况。用于实验的模型是:

所有实验均使用 Shadeform 的 GPU 市场完成,因此每次实验只需要花费几美元 H100。

实验结果表明,内存需求随着模型大小和训练方式的不同而显著变化。例如,全参数微调比 PEFT 需要更多的内存。

为什么 GRPO 对内存需求较高

这要从 GRPO 的原理说起,这是它的流程图。

GRPO 对内存需求较高的原因在于,其内部涉及多个模型,并且在训练数据中每个查询会产生多个输出。上图中的策略模型、参考模型和奖励模型各自都是一个需要进行推理的 LLM。(尽管从技术上讲,奖励模型可能不需要参数化,可以只是一个 Python 函数或正则表达式,但不影响 GRPO 对内存的高需求。)

为什么 8-Bit 优化和梯度检查点有助于减少内存占用?

通常来讲,训练一个大型语言模型需要在内存中存储三种主要类型的信息:模型参数、模型学习所需的梯度、优化器的跟踪数据。

对上述内容我们可以这样理解:如果模型的参数占用了 X 的空间,那么梯度也会占用大约相同的空间。然后,像 AdamW 这样的优化器需要更多的空间,因为它们就像一个记录员,跟踪最近的更新历史,以便更好地决定未来的优化。

为了减轻这种内存负担,通常采用两种技术:

  • 首先,可以使用像 AdamW 这样的 8-bit 优化器版本,它们能更高效地存储跟踪数据,同时仍保持良好的性能 —— 类似于压缩照片可以节省空间,同时保留大部分图像质量;

  • 其次,使用梯度检查点技术,这就像在训练过程中拍摄快照,而不是记录所有内容。虽然这会使训练速度减慢约 20-30%,但它显著减少了内存使用。

结合这些技术,即使对 GPU 资源有限的人来说,也能够训练更大的模型。

代码示例

像 trl 这样的库已经开始支持 GRPO,使得微调由 transformers 构成的 LLM 变得非常简单。代码也非常简洁,只需将训练器替换为 GRPOTrainer 并定义一些奖励即可。GRPO 的最小代码量大约只有 99 行,如果你使用的是像 meta-llama/Llama-3.2-1B-Instruct 这样的小型模型和像 openai/GSM8K 这样的数据集,可以非常快速地启动。

trl 项目地址:https://github.com/huggingface/trl?ref=ghost.oxen.ai

  1. import

  2. torch

  3. from

  4. datasets

  5. import

  6. load_dataset

  7. ,


  8. Dataset

  9. from

  10. transformers

  11. import


  12. AutoTokenizer

  13. ,


  14. AutoModelForCausalLM

  15. from

  16. trl

  17. import


  18. GRPOConfig

  19. ,


  20. GRPOTrainer

  21. import

  22. re

  23. SYSTEM_PROMPT

  24. =


  25. """

  26. Respond in the following format:

  27. <reasoning>

  28. ...

  29. </reasoning>

  30. <answer>

  31. ...

  32. </answer>

  33. """

  34. def

  35. extract_hash_answer

  36. (

  37. text

  38. :

  39. str

  40. )


  41. ->

  42. str

  43. |


  44. None

  45. :


  46. if


  47. "####"


  48. not


  49. in

  50. text

  51. :


  52. return


  53. None


  54. return

  55. text

  56. .

  57. split

  58. (

  59. "####"

  60. )[

  61. 1

  62. ].

  63. strip

  64. ()

  65. def

  66. get_gsm8k_questions

  67. (

  68. split

  69. =


  70. "train"

  71. )


  72. ->


  73. Dataset

  74. :

  75. data

  76. =

  77. load_dataset

  78. (

  79. 'openai/gsm8k'

  80. ,


  81. 'main'

  82. )[

  83. split

  84. ]

  85. data

  86. =

  87. data

  88. .

  89. map

  90. (

  91. lambda

  92. x

  93. :


  94. {


  95. 'prompt'

  96. :


  97. [


  98. {

  99. 'role'

  100. :


  101. 'system'

  102. ,


  103. 'content'

  104. :

  105. SYSTEM_PROMPT

  106. },


  107. {

  108. 'role'

  109. :


  110. 'user'

  111. ,


  112. 'content'

  113. :

  114. x

  115. [

  116. 'question'

  117. ]}


  118. ],


  119. 'answer'

  120. :

  121. extract_hash_answer

  122. (

  123. x

  124. [

  125. 'answer'

  126. ])


  127. })


  128. return

  129. data

  130. def

  131. extract_xml_answer

  132. (

  133. text

  134. :

  135. str

  136. )


  137. ->

  138. str

  139. :

  140. answer

  141. =

  142. text

  143. .

  144. split

  145. (

  146. "<answer>"

  147. )[-

  148. 1

  149. ]

  150. answer

  151. =

  152. answer

  153. .

  154. split

  155. (

  156. "</answer>"

  157. )[

  158. 0

  159. ]


  160. return

  161. answer

  162. .

  163. strip

  164. ()

  165. def

  166. format_reward_func

  167. (

  168. completions

  169. ,


  170. **

  171. kwargs

  172. )


  173. ->

  174. list

  175. [

  176. float

  177. ]:


  178. """Reward function that checks if the completion has a specific format."""

  179. pattern

  180. =

  181. r

  182. "^<reasoning>\n.*?\n</reasoning>\n<answer>\n.*?\n</answer>\n$"

  183. responses

  184. =


  185. [

  186. completion

  187. [

  188. 0

  189. ][

  190. "content"

  191. ]


  192. for

  193. completion

  194. in

  195. completions

  196. ]

  197. matches

  198. =


  199. [

  200. re

  201. .

  202. match

  203. (

  204. pattern

  205. ,

  206. r

  207. )


  208. for

  209. r

  210. in

  211. responses

  212. ]


  213. return


  214. [

  215. 0.5


  216. if

  217. match

  218. else


  219. 0.0


  220. for

  221. match

  222. in

  223. matches

  224. ]

  225. def

  226. accuracy_reward_func

  227. (

  228. prompts

  229. ,

  230. completions

  231. ,

  232. answer

  233. ,


  234. **

  235. kwargs

  236. )


  237. ->

  238. list

  239. [

  240. float

  241. ]:


  242. """Reward function that extracts the answer from the xml tags and compares it to the correct answer."""

  243. responses

  244. =


  245. [

  246. completion

  247. [

  248. 0

  249. ][

  250. 'content'

  251. ]


  252. for

  253. completion

  254. in

  255. completions

  256. ]

  257. extracted_responses

  258. =


  259. [

  260. extract_xml_answer

  261. (

  262. r

  263. )


  264. for

  265. r

  266. in

  267. responses

  268. ]


  269. return


  270. [

  271. 2.0


  272. if

  273. r

  274. ==

  275. a

  276. else


  277. 0.0


  278. for

  279. r

  280. ,

  281. a

  282. in

  283. zip

  284. (

  285. extracted_responses

  286. ,

  287. answer

  288. )]

  289. def

  290. main

  291. ():

  292. dataset

  293. =

  294. get_gsm8k_questions

  295. ()

  296. model_name

  297. =


  298. "meta-llama/Llama-3.2-1B-Instruct"

  299. model

  300. =


  301. AutoModelForCausalLM

  302. .

  303. from_pretrained

  304. (

  305. model_name

  306. ,

  307. torch_dtype

  308. =

  309. torch

  310. .

  311. bfloat16

  312. ,

  313. attn_implementation

  314. =

  315. "flash_attention_2"

  316. ,

  317. device_map

  318. =

  319. None


  320. ).

  321. to

  322. (

  323. "cuda"

  324. )

  325. tokenizer

  326. =


  327. AutoTokenizer

  328. .

  329. from_pretrained

  330. (

  331. model_name

  332. )

  333. tokenizer

  334. .

  335. pad_token

  336. =

  337. tokenizer

  338. .

  339. eos_token

  340. training_args

  341. =


  342. GRPOConfig

  343. (

  344. output_dir

  345. =

  346. "output"

  347. ,

  348. learning_rate

  349. =

  350. 5e-6

  351. ,

  352. adam_beta1

  353. =

  354. 0.9

  355. ,

  356. adam_beta2

  357. =

  358. 0.99

  359. ,

  360. weight_decay

  361. =

  362. 0.1

  363. ,

  364. warmup_ratio

  365. =

  366. 0.1

  367. ,

  368. lr_scheduler_type

  369. =

  370. 'cosine'

  371. ,

  372. logging_steps

  373. =

  374. 1

  375. ,

  376. bf16

  377. =

  378. True

  379. ,

  380. per_device_train_batch_size

  381. =

  382. 1

  383. ,

  384. gradient_accumulation_steps

  385. =

  386. 4

  387. ,

  388. num_generations

  389. =

  390. 4

  391. ,

  392. max_prompt_length

  393. =

  394. 256

  395. ,

  396. max_completion_length

  397. =

  398. 786

  399. ,

  400. num_train_epochs

  401. =

  402. 1

  403. ,

  404. save_steps

  405. =

  406. 100

  407. ,

  408. save_total_limit

  409. =

  410. 1

  411. ,

  412. max_grad_norm

  413. =

  414. 0.1

  415. ,

  416. log_on_each_node

  417. =

  418. False

  419. ,


  420. )

  421. trainer

  422. =


  423. GRPOTrainer

  424. (

  425. model

  426. =

  427. model

  428. ,

  429. processing_class

  430. =

  431. tokenizer

  432. ,

  433. reward_funcs

  434. =[

  435. format_reward_func

  436. ,

  437. accuracy_reward_func


  438. ],

  439. args

  440. =

  441. training_args

  442. ,

  443. train_dataset

  444. =

  445. dataset

  446. ,


  447. )

  448. trainer

  449. .

  450. train

  451. ()

  452. if

  453. __name__

  454. ==


  455. "__main__"

  456. :

  457. main

  458. ()

Num Generations 有什么用

Num Generations 是一个超参数,它决定了我们将在训练数据中对每个查询采样多少个补全。然而,这会显著增加 VRAM 的消耗。

目前有一个开放的 GitHub 问题,可能会帮助解决内存瓶颈问题,可以参考如下链接

地址:https://github.com/huggingface/trl/issues/2709?ref=ghost.oxen.ai

对于 num_completions=8,16,64 (DeepSeekMath 论文使用的 64),作者表示,不用再次计算上述所有值,而是使用了 1B 参数模型进行了测试,以显示内存增长。不过,作者还是建议大家在内存瓶颈得到修复之前使用 num_generations=4,也能获得不错的性能。

影响 VRAM 的一些因素

要对所有影响显存(VRAM)使用的因素进行全面的超参数验证,需要进行大量的实验。简单起见,这里只指出了需要注意的设置,以及实验中使用的具体数值。

  • batch_size=1,由于 GRPO 为每个查询生成多个响应,batch size 会迅速失控。

  • gradient_accumulation_steps=4,优化器是另一个占用大量 VRAM 的地方。此参数决定了我们将存储的梯度以帮助优化器进行其「爬山」过程。

  • num_completions=4,DeepSeekMath 论文中使用了 64。这完全超出了有些人的计算预算。

  • max_prompt_length=256,如果你想训练模型拥有更大上下文的推理能力,将不得不增加 VRAM。GSM8K 的提示相对较小,适合此测试。

  • max_completion_length=786,同样,由于计算注意力的内存有限,推理链在这里受到限制。上下文或生成的 token 越多,需要的内存就越大。

  • LoRA target_modules=["q_proj", "k_proj", "o_proj", "up_proj", "down_proj"] 在这方面可以尝试几种不同的迭代。target_modules="all-linear" 是一种流行的方式,可以从你的 LoRA 中挤出最多的性能(就准确性而言)。

对 VRAM 使用的粗略估算

如果你正在使用 FP16 精度进行训练,以下是一些简单的估算方法,可以帮助你了解内存主要用在了哪些地方:

  • 模型参数:每个参数占用 2 字节。

  • 参考模型参数:每个参数占用 2 字节。

  • 梯度:每个参数占用 2 字节。

  • 优化器状态:每个参数占用 8 字节。

  • 8 位优化器:每个参数占用 4 字节。

  • PEFT:有助于减少梯度的显存占用。

最后是关于准确率的。作者完成了一个 10 亿参数的 Llama 3.2 模型的完整训练。在应用 GRPO 之前,该模型在保留测试集上达到了约 19% 的准确率,而在经过一个训练周期后,模型的准确率飙升至约 40.5%。虽然这离 SOTA 水平还差得很远,但这展示了 GRPO 的强大潜力。

举报/反馈
分享到: 微博 QQ 空间
对本文内容有合作意向?
我们将在 1 个工作日内与您联系
留言咨询