DeepSeek用的GRPO占用大量内存?有人给出了些破解方法
大家好,今天本人给大家带来文章《DeepSeek用的GRPO占用大量内存?有人给出了些破解方法》,文中内容主要涉及到,如果你对科技周边方面的知识点感兴趣,那就请各位朋友继续看下去吧~希望能真正帮到你们,谢谢!
RTX 3080 移动版训练大型语言模型的实用指南
本文旨在指导 GPU 资源受限的开发者如何利用 GRPO (Group Relative Policy Optimization) 训练大型语言模型。DeepSeek-R1 的发布使得 GRPO 成为强化学习训练大型语言模型的热门方法,因为它高效且易于训练。 GRPO 通过利用模型自身生成的训练数据进行迭代改进,目标是最大化生成文本的优势函数,同时保持模型与参考策略的接近性。

选择合适的模型大小和训练方法(全参数微调或参数高效微调 - PEFT)是训练的关键。本文作者 Greg Schoeninger (Oxen.ai CEO) 使用配备 16GB 显存的 RTX 3080 笔记本电脑进行实验,并分享了其经验。
原文链接:https://www.oxen.ai/blog/grpo-vram-requirements-for-the-gpu-poor{{TABLE_PLACEHOLDER_1}}
作者在使用 trl 库的 GRPO 实现时,遇到了显存不足 (OOM) 错误:
torch.OutOfMemoryError: CUDA out of memory.
Tried to allocate 1.90 GiB. GPU 0 has a total capacity of 15.73 GiB of which 1.28 GiB is free.
Including non-PyTorch memory, this process has 14.43 GiB memory in use. Of the allocated memory 11.82 GiB is allocated by PyTorch, and 2.41 GiB is reserved by PyTorch but unallocated. If reserved but unallocated memory is large try setting PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True to avoid fragmentation. See documentation for Memory Management (https://pytorch.org/docs/stable/notes/cuda.html#environment-variables)
实验结果与内存需求分析
作者进行了一系列实验,测试不同模型大小(5亿到140亿参数)在 GSM8K 数据集上训练前 100 步的峰值内存使用情况,并比较了全参数微调和 PEFT 的内存需求。所有实验均在 Nvidia H100 上完成。

使用的模型包括:

GRPO 对内存需求高的原因在于其内部涉及多个模型(策略模型、参考模型、奖励模型)以及每个查询产生的多个输出。

优化内存使用
8位优化器和梯度检查点技术可以有效减少内存占用。8位优化器更高效地存储优化器状态,而梯度检查点则通过在训练过程中拍摄快照来减少内存使用,虽然会降低训练速度。
代码示例
trl 库简化了 GRPO 的使用。以下代码示例展示了如何使用 trl 训练小型模型:
import torch
from datasets import load_dataset, Dataset
from transformers import AutoTokenizer, AutoModelForCausalLM
from trl import GRPOConfig, GRPOTrainer
import re
SYSTEM_PROMPT = """
Respond in the following format:
...
...
"""
def extract_hash_answer(text: str) -> str | None:
if "####" not in text:
return None
return text.split("####")[1].strip()
def get_gsm8k_questions(split = "train") -> Dataset:
data = load_dataset('openai/gsm8k', 'main')[split]
data = data.map(lambda x: {
'prompt': [
{'role': 'system', 'content': SYSTEM_PROMPT},
{'role': 'user', 'content': x['question']}
],
'answer': extract_hash_answer(x['answer'])
})
return data
def extract_xml_answer(text: str) -> str:
answer = text.split("")[-1]
answer = answer.split("")[0]
return answer.strip()
def format_reward_func(completions, **kwargs) -> list[float]:
"""Reward function that checks if the completion has a specific format."""
pattern = r"^\n.*?\n \n\n.*?\n \n$"
responses = [completion[0]["content"] for completion in completions]
matches = [re.match(pattern, r) for r in responses]
return [0.5 if match else 0.0 for match in matches]
def accuracy_reward_func(prompts, completions, answer, **kwargs) -> list[float]:
"""Reward function that extracts the answer from the xml tags and compares it to the correct answer."""
responses = [completion[0]['content'] for completion in completions]
extracted_responses = [extract_xml_answer(r) for r in responses]
return [2.0 if r == a else 0.0 for r, a in zip(extracted_responses, answer)]
def main():
dataset = get_gsm8k_questions()
model_name = "meta-llama/Llama-3.2-1B-Instruct"
model = AutoModelForCausalLM.from_pretrained(
model_name,
torch_dtype=torch.bfloat16,
attn_implementation="flash_attention_2",
device_map=None
).to("cuda")
tokenizer = AutoTokenizer.from_pretrained(model_name)
tokenizer.pad_token = tokenizer.eos_token
training_args = GRPOConfig(
output_dir="output",
learning_rate=5e-6,
adam_beta1=0.9,
adam_beta2=0.99,
weight_decay=0.1,
warmup_ratio=0.1,
lr_scheduler_type='cosine',
logging_steps=1,
bf16=True,
per_device_train_batch_size=1,
gradient_accumulation_steps=4,
num_generations=4,
max_prompt_length=256,
max_completion_length=786,
num_train_epochs=1,
save_steps=100,
save_total_limit=1,
max_grad_norm=0.1,
log_on_each_node=False,
)
trainer = GRPOTrainer(
model=model,
processing_class=tokenizer,
reward_funcs=[
format_reward_func,
accuracy_reward_func
],
args=training_args,
train_dataset=dataset,
)
trainer.train()
if __name__ == "__main__":
main()
trl 项目地址:https://github.com/huggingface/trl?ref=ghost.oxen.ai
超参数调整与VRAM估算
num_generations 超参数会显著影响 VRAM 消耗。建议在内存瓶颈解决前使用 num_generations=4。

GitHub 问题讨论:https://github.com/huggingface/trl/issues/2709?ref=ghost.oxen.ai
其他影响 VRAM 的因素包括 batch_size、gradient_accumulation_steps、max_prompt_length、max_completion_length 和 LoRA 的 target_modules。

最后,作者分享了 10 亿参数 Llama 3.2 模型的训练结果,展示了 GRPO 在提高模型准确率方面的潜力。
通过合理的参数设置和优化技术,即使使用资源有限的 RTX 3080 移动版 GPU,也能有效训练大型语言模型。
到这里,我们也就讲完了《DeepSeek用的GRPO占用大量内存?有人给出了些破解方法》的内容了。个人认为,基础知识的学习和巩固,是为了更好的将其运用到项目中,欢迎关注golang学习网公众号,带你了解更多关于产业,GRPO的知识点!
腾势Z9迎OTA重磅升级,全量推送城市领航功能
- 上一篇
- 腾势Z9迎OTA重磅升级,全量推送城市领航功能
- 下一篇
- TIL:使用ModuleCreateRequire(节点)在ES模型中同步导入
-
- 科技周边 · 人工智能 | 2小时前 |
- RAG 文档切块按标题层级保留语义边界
- 300浏览 收藏
-
- 科技周边 · 人工智能 | 14小时前 | 人工智能 · rag · 语义检索重排器 第二阶段重排 Cross-Encoder Retrieve and Re-Rank Recall@K
- 语义检索重排器何时值得加入第二阶段
- 449浏览 收藏
-
- 科技周边 · 人工智能 | 18小时前 | 缓存 · 人工智能 · 提示词工程 · 提示词缓存 cache_control Prompt Caching 静态前缀 cache_read_input_tokens
- 提示词缓存命中率低应如何划分静态前缀
- 453浏览 收藏
-
- 科技周边 · 人工智能 | 21小时前 |
- 模型路由怎样按任务难度分配不同推理预算
- 232浏览 收藏
-
- 科技周边 · 人工智能 | 23小时前 |
- 结构化输出遇到递归字段时怎样约束模式
- 215浏览 收藏
-
- 科技周边 · 人工智能 | 1天前 | 人工智能 · 工具调用 ·
- 智能体工具调用失败后怎样设计可控重试
- 236浏览 收藏
-
- 科技周边 · 人工智能 | 1天前 |
- 多模态模型输入图片过大时如何控制视觉令牌
- 196浏览 收藏
-
- 科技周边 · 人工智能 | 1天前 |
- 推理服务连续批处理怎样减少 GPU 空转
- 404浏览 收藏
-
- 科技周边 · 人工智能 | 1天前 | 人工智能 · LoRa PEFT merge_and_unload 量化推理 模型合并
- LoRA 合并权重后输出变化过大应检查什么
- 404浏览 收藏
-
- 科技周边 · 人工智能 | 1天前 |
- RAG 分块重叠过大为什么会降低检索多样性
- 409浏览 收藏
-
- 前端进阶之JavaScript设计模式
- 设计模式是开发人员在软件开发过程中面临一般问题时的解决方案,代表了最佳的实践。本课程的主打内容包括JS常见设计模式以及具体应用场景,打造一站式知识长龙服务,适合有JS基础的同学学习。
- 543次学习
-
- GO语言核心编程课程
- 本课程采用真实案例,全面具体可落地,从理论到实践,一步一步将GO核心编程技术、编程思想、底层实现融会贯通,使学习者贴近时代脉搏,做IT互联网时代的弄潮儿。
- 516次学习
-
- 简单聊聊mysql8与网络通信
- 如有问题加微信:Le-studyg;在课程中,我们将首先介绍MySQL8的新特性,包括性能优化、安全增强、新数据类型等,帮助学生快速熟悉MySQL8的最新功能。接着,我们将深入解析MySQL的网络通信机制,包括协议、连接管理、数据传输等,让
- 500次学习
-
- JavaScript正则表达式基础与实战
- 在任何一门编程语言中,正则表达式,都是一项重要的知识,它提供了高效的字符串匹配与捕获机制,可以极大的简化程序设计。
- 487次学习
-
- 从零制作响应式网站—Grid布局
- 本系列教程将展示从零制作一个假想的网络科技公司官网,分为导航,轮播,关于我们,成功案例,服务流程,团队介绍,数据部分,公司动态,底部信息等内容区块。网站整体采用CSSGrid布局,支持响应式,有流畅过渡和展现动画。
- 485次学习
-
- PubMedQA
- 深入了解PubMedQA生物医学问答数据集,涵盖其核心功能、使用方法及在临床决策、药物研发等场景的应用,助力提升NLP模型性能。
- 402次使用
-
- H2O EvalGPT
- H2O EvalGPT是H2O.ai推出的开源LLM评估平台,提供详细的大模型性能排行榜、行业特定基准测试及A/B测试功能,助您快速选择最适合项目的高性能大语言模型。
- 478次使用
-
- LMArena
- LMArena是加州大学伯克利分校推出的AI模型匿名评测平台。通过盲测投票机制,用户可对比不同大模型回答并生成实时排行榜,助力开发者优化模型及用户选择最佳AI工具。
- 487次使用
-
- HELM
- 深入了解斯坦福推出的HELM(Holistic Evaluation of Language Models)大模型评测体系。本文解析其核心功能、安装配置步骤及应用场景,涵盖准确性、公平性、鲁棒性等多维度指标,助力开发者全面优化语言模型性能。
- 435次使用
-
- MMBench
- MMBench是由上海人工智能实验室等机构联合推出的多模态基准测试平台,提供细粒度能力评估、大规模数据集及VLMEvalKit工具。本文详细介绍其核心功能、安装使用方法及应用场景,助力开发者全面评估多模态模型性能。
- 260次使用
-
- 本地大模型反复输出同一句话怎么调整生成参数
- 2026-09-06 501浏览
-
- Python 调用大模型时如何用结构化输出校验 JSON:从解析失败到可重试
- 2026-08-29 501浏览
-
- AI写作工具免费版安装教程(含豆包Clawdbot)
- 2026-05-30 501浏览
-
- WPS AI能自动生成PPT吗?输入主题一键制作演示文稿
- 2026-05-27 501浏览
-
- Canva手机闪退解决方法及适配指南
- 2026-05-25 501浏览

