CPM-Distill:经过知识蒸馏的小型文本生成模型

本文介绍知识蒸馏技术及基于PaddleNLP加载CPM-Distill模型实现文本生成。知识蒸馏是模型压缩方法,以“教师-学生网络”思想,让简单模型拟合复杂模型输出,效果优于从头训练。CPM-Distill由GPT-2 Large蒸馏得到,文中还给出安装依赖、加载模型、解码方法及文本生成示例。

☞☞☞AI 智能聊天, 问答助手, AI 智能搜索, 免费无限量使用 DeepSeek R1 模型☜☜☜

cpm-distill:经过知识蒸馏的小型文本生成模型 - 创想鸟

引入

近些年来,随着 Bert 这样的大规模预训练模型的问世,NLP 领域的模型也逐渐变得越来越大了受限于算力水平,如此大规模的模型要应用在实际的部署场景都是不太实际的因此需要通过一些方式对大规模的模型进行压缩,使其能够在部署场景下达到一个相对可用的速度常见的模型压缩方法有:剪枝、量化、知识蒸馏等最近 CPM(Chinese Pre-Trained Models)项目又开源了一个使用知识蒸馏得到的小型文本生成模型 CPM-Distill本次项目就简单介绍一下知识蒸馏技术并且通过 PaddleNLP 套件加载 CPM-Distill 模型实现文本生成

相关项目

Paddle2.0:构建一个经典的文本生成模型GPT-2文本生成:使用GPT-2加载CPM-LM模型实现简单的问答机器人文本生成:让AI帮你写文章吧【AI创造营】PaddleHub 配合 PaddleNLP 实现简单的文本生成

相关资料

论文:CPM: A Large-scale Generative Chinese Pre-trained Language ModelDistilling the Knowledge in a Neural Network官方实现:TsinghuaAI/CPM-Distill

模型压缩技术

CPM-Distill:经过知识蒸馏的小型文本生成模型 - 创想鸟

知识蒸馏(Knowledge Distillation)

知识蒸馏是一种模型压缩方法,是一种基于“教师-学生网络思想”的训练方法。

由 Hinton 在 2015 年 Distilling the Knowledge in a Neural Network 的论文首次提出了知识蒸馏的并尝试在 CV 领域中使用,旨在把大模型学到的知识灌输到小模型中,以达到缩小模型的目标,示意图如下:

CPM-Distill:经过知识蒸馏的小型文本生成模型 - 创想鸟

说人话就是指用一个简单模型去拟合复杂模型的输出,这个输出也叫做“软标签”,当然也可以加入真实数据作为“硬标签”一同训练。使用知识蒸馏技术相比直接从头训练的效果一般会更好一些,因为教师模型能够指导学生模型收敛到一个更佳的位置。

CPM-Distill:经过知识蒸馏的小型文本生成模型 - 创想鸟

知识蒸馏技术除了可以用来将网络从大网络转化成一个小网络,并保留接近于大网络的性能;也可以将多个网络的学到的知识转移到一个网络中,使得单个网络的性能接近 emsemble 的结果。

蒸馏模型信息

教师模型为 GPT-2 Large,具体的模型参数如下:

teacher_model = GPTModel(    vocab_size=30000,    hidden_size=2560,    num_hidden_layers=32,    num_attention_heads=32,    intermediate_size=10240,    hidden_act="gelu",    hidden_dropout_prob=0.1,    attention_probs_dropout_prob=0.1,    max_position_embeddings=1024,    type_vocab_size=1,    initializer_range=0.02,    pad_token_id=0,    topo=None)

学生模型为 GPT-2 Small,具体的模型参数如下:

teacher_model = GPTModel(    vocab_size=30000,    hidden_size=768,    num_hidden_layers=12,    num_attention_heads=12,    intermediate_size=3072,    hidden_act="gelu",    hidden_dropout_prob=0.1,    attention_probs_dropout_prob=0.1,    max_position_embeddings=1024,    type_vocab_size=1,    initializer_range=0.02,    pad_token_id=0,    topo=None)

蒸馏 loss

将大模型和小模型每个位置上输出之间的 KL 散度作为蒸馏 loss,同时加上原来的 language model loss。总 loss 如下:

CPM-Distill:经过知识蒸馏的小型文本生成模型 - 创想鸟

其中 LlmLlm 为 GPT-2 原始的 language modeling loss。

安装依赖

In [ ]

!pip install paddlenlp==2.0.1 sentencepiece==0.1.92

加载模型

In [1]

import paddlefrom paddlenlp.transformers import GPTModel, GPTForPretraining, GPTChineseTokenizer# tokenizer 与 CPM-LM 模型一致tokenizer = GPTChineseTokenizer.from_pretrained('gpt-cpm-large-cn')# 实例化 GPT2-small 模型gpt = GPTModel(    vocab_size=30000,    hidden_size=768,    num_hidden_layers=12,    num_attention_heads=12,    intermediate_size=3072,    hidden_act="gelu",    hidden_dropout_prob=0.1,    attention_probs_dropout_prob=0.1,    max_position_embeddings=1024,    type_vocab_size=1,    initializer_range=0.02,    pad_token_id=0,    topo=None)# 加载预训练模型参数params = paddle.load('data/data92160/gpt-cpm-small-cn-distill.pdparams')# 设置参数gpt.set_dict(params)# 使用 GPTForPretraining 向模型中添加输出层model = GPTForPretraining(gpt)# 将模型设置为评估模式model.eval()
[2021-05-28 19:38:04,469] [    INFO] - Found /home/aistudio/.paddlenlp/models/gpt-cpm-large-cn/gpt-cpm-cn-sentencepiece.model

模型解码

In [40]

import paddleimport numpy as np# Greedy Searchdef greedy_search(text, max_len=32, end_word=None):    # # 终止标志    if end_word is not None:        stop_id = tokenizer.encode(end_word)['input_ids']        length = len(stop_id)    else:        stop_id = [tokenizer.eod_token_id]        length = len(stop_id)        # 初始预测    ids = tokenizer.encode(text)['input_ids']    input_id = paddle.to_tensor(np.array(ids).reshape(1, -1).astype('int64'))    output, cached_kvs = model(input_id, use_cache=True)    next_token = int(np.argmax(output[0, -1].numpy()))    ids.append(next_token)    # 使用缓存进行继续预测    for i in range(max_len-1):        input_id = paddle.to_tensor(np.array([next_token]).reshape(1, -1).astype('int64'))        output, cached_kvs = model(input_id, use_cache=True, cache=cached_kvs)        next_token = int(np.argmax(output[0, -1].numpy()))        ids.append(next_token)        # 根据终止标志停止预测        if ids[-length:]==stop_id:            if end_word is None:               ids = ids[:-1]            break        return tokenizer.convert_ids_to_string(ids)

In [39]

import paddleimport numpy as np# top_k and top_p filteringdef top_k_top_p_filtering(logits, top_k=0, top_p=1.0, filter_value=-float('Inf')):    """ Filter a distribution of logits using top-k and/or nucleus (top-p) filtering        Args:            logits: logits distribution shape (vocabulary size)            top_k > 0: keep only top k tokens with highest probability (top-k filtering).            top_p > 0.0: keep the top tokens with cumulative probability >= top_p (nucleus filtering).                Nucleus filtering is described in Holtzman et al. (http://arxiv.org/abs/1904.09751)        From: https://gist.github.com/thomwolf/1a5a29f6962089e871b94cbd09daf317    """    top_k = min(top_k, logits.shape[-1])  # Safety check    logits_np = logits.numpy()    if top_k > 0:        # Remove all tokens with a probability less than the last token of the top-k        indices_to_remove = logits_np < np.sort(logits_np)[-top_k]        logits_np[indices_to_remove] = filter_value    if top_p  top_p        # Shift the indices to the right to keep also the first token above the threshold        sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1]        sorted_indices_to_remove[..., 0] = 0        indices_to_remove = sorted_indices[sorted_indices_to_remove]        logits_np[indices_to_remove] = filter_value    return paddle.to_tensor(logits_np)# Nucleus Sampledef nucleus_sample(text, max_len=32, end_word=None, repitition_penalty=1.0, temperature=1.0, top_k=0, top_p=1.0):    # 终止标志    if end_word is not None:        stop_id = tokenizer.encode(end_word)['input_ids']        length = len(stop_id)    else:        stop_id = [tokenizer.eod_token_id]        length = len(stop_id)    # 初始预测    ids = tokenizer.encode(text)['input_ids']    input_id = paddle.to_tensor(np.array(ids).reshape(1, -1).astype('int64'))    output, cached_kvs = model(input_id, use_cache=True)    next_token_logits = output[0, -1, :]    for id in set(ids):        next_token_logits[id] /= repitition_penalty    next_token_logits = next_token_logits / temperature    filtered_logits = top_k_top_p_filtering(next_token_logits, top_k=top_k, top_p=top_p)    next_token = paddle.multinomial(paddle.nn.functional.softmax(filtered_logits, axis=-1), num_samples=1).numpy()    ids += [int(next_token)]    # 使用缓存进行继续预测    for i in range(max_len-1):        input_id = paddle.to_tensor(np.array([next_token]).reshape(1, -1).astype('int64'))        output, cached_kvs = model(input_id, use_cache=True, cache=cached_kvs)        next_token_logits = output[0, -1, :]        for id in set(ids):            next_token_logits[id] /= repitition_penalty        next_token_logits = next_token_logits / temperature        filtered_logits = top_k_top_p_filtering(next_token_logits, top_k=top_k, top_p=top_p)        next_token = paddle.multinomial(paddle.nn.functional.softmax(filtered_logits, axis=-1), num_samples=1).numpy()        ids += [int(next_token)]        # 根据终止标志停止预测        if ids[-length:]==stop_id:            if end_word is None:               ids = ids[:-1]            break    return tokenizer.convert_ids_to_string(ids)

文本生成

In [41]

# 输入文本inputs = input('请输入文本:')print(inputs)# 使用 Nucleus Sample 进行文本生成outputs = greedy_search(    inputs, # 输入文本    max_len=128, # 最大生成文本的长度    end_word=None)# 打印输出print(outputs)
请输入文本:请在此处输入你的姓名请在此处输入你的姓名,然后点击“确定”,就可以开始游戏了。游戏目标:在限定时间内,成功地把所有的牌都通通打完。

In [43]

# 输入文本inputs = input('请输入文本:')print(inputs)for x in range(5):    # 使用 Nucleus Sample 进行文本生成    outputs = nucleus_sample(        inputs, # 输入文本        max_len=128, # 最大生成文本的长度        end_word='。', # 终止符号        repitition_penalty=1.0, # 重复度抑制        temperature=1.0, # 温度        top_k=3000, # 取前k个最大输出再进行采样        top_p=0.9 # 抑制概率低于top_p的输出再进行采样    )    # 打印输出    print(outputs)
请输入文本:请在此处输入你的姓名请在此处输入你的姓名、学校、专业及学科,并在社交媒体上公布你的个人简介。请在此处输入你的姓名或者电话,对方会及时通知你。请在此处输入你的姓名、民族及籍贯信息,当您找到 CADULI 的联系方式后,我们会按您所选择的申请中心,以电子邮件的形式向您发送邮件。请在此处输入你的姓名和电话号码,由资深会所接待员进行介绍,因为此处有不少中国的大老板,英文能看。请在此处输入你的姓名、联系电话、银行卡号和手机号。

以上就是CPM-Distill:经过知识蒸馏的小型文本生成模型的详细内容,更多请关注创想鸟其它相关文章!

版权声明:本文内容由互联网用户自发贡献,该文观点仅代表作者本人。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。
如发现本站有涉嫌抄袭侵权/违法违规的内容, 请发送邮件至 chuangxiangniao@163.com 举报,一经查实,本站将立刻删除。
发布者:程序猿,转转请注明出处:https://www.chuangxiangniao.com/p/66523.html

(0)
打赏 微信扫一扫 微信扫一扫 支付宝扫一扫 支付宝扫一扫
《原神》重磅UGC玩法官宣!玩家能自制Moba、射击、塔防等多种玩法
上一篇 2025年11月12日 07:17:23
一本漫画app热门连载_一本漫画软件直达入口
下一篇 2025年11月12日 07:20:27

相关推荐

  • Windows11内存占用率过高怎么解决_Windows11内存占用过高修复方法

    1、通过任务管理器结束高内存占用进程;2、禁用Superfetch(SysMain)服务以降低内存负担;3、优化启动项减少后台负载;4、升级物理内存条提升系统性能。 如果您发现Windows 11系统运行缓慢,并且任务管理器显示内存占用率持续处于高位,这可能是由于后台进程过多、系统服务占用资源或硬件…

    2026年9月21日
    100
  • mysql常用存储引擎有哪些

    InnoDB是现代MySQL应用的首选存储引擎,因其支持事务(ACID)、行级锁、外键约束、崩溃恢复和MVCC,适用于高并发、数据完整性要求高的OLTP场景;MyISAM虽读取快但仅支持表级锁且无事务和外键,适用于读多写少的简单场景,已逐渐被淘汰;Memory引擎将数据存于内存,速度快但易失,适合临…

    2026年9月21日
    000
  • 利用蝴蝶号搭建多账号无人直播系统的完整方案

    利用蝴蝶号搭建多账号无人直播系统的完整方案利用蝴蝶号搭建多账号无人直播系统的完整方案利用蝴蝶号搭建多账号无人直播系统的完整方案利用蝴蝶号搭建多账号无人直播系统的完整方案

    搭建多账号无人直播系统并非一键操作,而是通过“蝴蝶号”实现自动化流程。首先,“蝴蝶号”负责多账号的生命周期管理,包括登录、状态维护、ip代理分配和设备指纹模拟;其次,内容调度系统决定直播内容及播放时间,可为预录视频或动态生成流;再次,推流引擎将内容实时推送至平台,推荐使用ffmpeg结合python…

    2026年9月21日 用户投稿
    100
  • 锚定AI终端存储市场,康盈半导体连发三款新品

    锚定AI终端存储市场,康盈半导体连发三款新品锚定AI终端存储市场,康盈半导体连发三款新品锚定AI终端存储市场,康盈半导体连发三款新品锚定AI终端存储市场,康盈半导体连发三款新品

    ☞☞☞AI 智能聊天, 问答助手, AI 智能搜索, 免费无限量使用 DeepSeek R1 模型☜☜☜ 三款新品聚焦AI存储需求 在最新举行的产品发布会上,康盈半导体正式推出三款专为AI应用场景打造的全新存储解决方案,覆盖嵌入式存储与高性能固态硬盘等多个品类,旨在满足多样化AI终端对高效、紧凑、低…

    2026年9月21日 用户投稿
    100
  • 数据库运维开发环境的调试模式演进

    数据库运维开发环境的调试模式演进数据库运维开发环境的调试模式演进数据库运维开发环境的调试模式演进数据库运维开发环境的调试模式演进

    这是学习笔记的第2393篇文章。 昨日,同事反馈了一个问题,原本的办公机环境中的虚拟机可以将办公机的IP暴露出来,提供数据库运维的API服务。例如,办公机的IP为192.168.10.100,而使用VirtualBox的虚拟机采用主机模式,其IP可能为192.168.56.100,那么192.168…

    2026年9月21日 用户投稿
    100
  • linux内核定时器实验

    linux内核定时器实验linux内核定时器实验linux内核定时器实验linux内核定时器实验

    大家好,又见面了,我是你们的朋友全栈君。 文章目录一、linux时间管理和内核定时器简介1.内核时间管理简介2.内核定时器简介1.init_timer 函数2.add_timer 函数3.del_timer 函数4.del_timer_sync 函数5.mod_timer 函数3.linux内核短延…

    2026年9月21日 用户投稿
    000
  • WordPress插件定制:使用Filter Hook修改邮件通知接收者

    本教程将指导您如何在WordPress中利用Filter Hook定制插件行为,特别是修改第三方插件的邮件通知接收者。我们将详细讲解如何识别目标Filter、理解其参数,并正确编写回调函数来拦截或修改数据,以实现自定义的邮件发送逻辑,避免因参数不匹配导致的错误。 WordPress Hook机制概览…

    2026年9月21日
    100
  • VSCode编写Java代码方法_VSCode搭建Java开发环境实战教程

    答案:在VSCode中配置Java开发环境需安装JDK并设置环境变量,再安装VSCode及Java扩展包,即可实现Java项目的创建、编写、运行与调试。它轻量、启动快,支持多语言和丰富扩展,集成Maven/Gradle,适合日常开发。 在VSCode里编写Java代码,说白了,就是把这个轻量级的代码…

    2026年9月21日
    100
  • JavaScript中的模块联邦如何实现微前端的代码共享?

    模块联邦通过运行时动态加载实现微前端代码共享,无需打包公共依赖。使用 ModuleFederationPlugin 配置 name、remotes、exposes 和 shared,使应用可暴露或引入远程模块,支持组件、工具函数及状态管理共享,提升复用性并减少冗余。 模块联邦通过在构建时让不同应用直…

    2026年9月21日
    200
  • Swoole如何实现一个UDP服务器

    答案:使用Swoole可轻松创建高性能UDP服务器。通过new SwooleServer()设置UDP套接字,监听Packet事件接收数据,利用sendto()回复客户端;结合set()配置worker_num等参数优化性能,配合PHP UDP客户端测试通信,适用于高并发、低延迟场景。 使用Swoo…

    2026年9月21日
    100
  • MySQL执行计划中的Extra字段代表什么_怎么看优化空间?

    MySQL执行计划中的Extra字段代表什么_怎么看优化空间?MySQL执行计划中的Extra字段代表什么_怎么看优化空间?MySQL执行计划中的Extra字段代表什么_怎么看优化空间?MySQL执行计划中的Extra字段代表什么_怎么看优化空间?

    在 mysql 查询优化中,执行计划的 extra 字段用于说明查询执行时的额外操作,常见的值包括:1. using filesort 表示需要额外排序,应尽量通过建立索引避免;2. using temporary 表示使用了临时表,常见于 group by 或复杂 join,需优化减少其使用;3.…

    2026年9月21日 用户投稿
    100
  • 如何通过tracert命令追踪数据包从本地到目标服务器的完整路径?

    打开命令提示符,输入cmd并回车;2. 执行tracert 目标地址命令追踪路径;3. 查看每跳响应时间与IP,分析延迟变化定位网络瓶颈;4. 注意部分节点可能因防火墙不响应导致超时。 使用 tracert(Windows 系统)命令可以追踪数据包从你的计算机到目标服务器所经过的每一跳网络节点,帮助…

    2026年9月21日
    1000
  • Linux interfaces 虚拟网络类型了解01

    Linux interfaces 虚拟网络类型了解01Linux interfaces 虚拟网络类型了解01Linux interfaces 虚拟网络类型了解01Linux interfaces 虚拟网络类型了解01

    在osi模型的定义中,数据链路层和物理层,以及传输层和网络层执行的任务在概念上相似:它们都提供了数据传输的方式,即沿着特定路径将数据从源点传输到目的地的方法。然而,数据链路层和物理层负责跨物理路径的通信服务,而传输层和网络层则提供由多个数据链路组成的逻辑路径或虚拟路径的通信服务。 Bridge操作指…

    2026年9月21日 用户投稿
    100
  • 如何在Java中理解Java I/O与NIO机制

    传统I/O是阻塞式流模型,适用于低并发场景;NIO基于缓冲区与通道,支持非阻塞和多路复用,适合高并发网络应用,核心区别在于线程模型与资源利用率。 Java中的I/O(输入/输出)与NIO(New I/O)是处理数据读写的核心机制,理解它们的区别和使用场景对开发高性能应用至关重要。传统I/O基于流模型…

    2026年9月21日
    100
  • JavaScript中的尾调用优化(TCO)在ES6中如何工作?

    尾调用是指函数的最后一个动作调用另一个函数,ES6引入尾调用优化以重用栈帧、避免内存溢出,支持真正的尾递归,如阶乘函数通过累积参数实现。 尾调用优化(Tail Call Optimization, TCO)是ES6引入的一项语言特性,目的是在特定条件下重用函数调用栈帧,避免不必要的内存增长,从而支持…

    2026年9月21日
    200
  • 抖音蝴蝶号无人直播带货操作流程及注意事项

    抖音蝴蝶号无人直播带货操作流程及注意事项抖音蝴蝶号无人直播带货操作流程及注意事项抖音蝴蝶号无人直播带货操作流程及注意事项抖音蝴蝶号无人直播带货操作流程及注意事项

    “抖音蝴蝶号无人直播带货”是一种通过自动化或半自动化技术实现的直播销售模式。①其核心在于摆脱真人主播限制,实现24小时不间断直播,提升效率与流量利用率;②关键步骤包括明确账号定位与商品选择、准备高质量且丰富的内容素材、利用虚拟人或预录内容实现直播推流、结合智能客服模拟评论区互动;③优势在于降低人力成…

    2026年9月21日 用户投稿
    600
  • VSCode侧边栏怎么去掉_VSCode侧边栏隐藏教程

    隐藏VSCode侧边栏可通过Ctrl + B(Windows/Linux)或Cmd + B(macOS)快捷键快速切换,也可通过菜单栏“视图 > 外观 > 切换侧边栏可见性”或命令面板执行“View: Toggle Sidebar Visibility”实现。推荐使用快捷键操作,效率最高…

    2026年9月21日
    100
  • MobileCLIP2— 苹果开源的端侧多模态模型

    MobileCLIP2— 苹果开源的端侧多模态模型MobileCLIP2— 苹果开源的端侧多模态模型MobileCLIP2— 苹果开源的端侧多模态模型MobileCLIP2— 苹果开源的端侧多模态模型

    ☞☞☞AI 智能聊天, 问答助手, AI 智能搜索, 免费无限量使用 DeepSeek R1 模型☜☜☜ 可图大模型 可图大模型(Kolors)是快手大模型团队自研打造的文生图AI大模型 32 查看详情 MobileCLIP2是什么 mobileclip2是由苹果研究团队开发的新一代高效多模态模型,…

    2026年9月21日 用户投稿
    300
  • MySQL用户权限体系配置思路_Sublime中编辑多用户分权管理脚本

    MySQL用户权限体系配置思路_Sublime中编辑多用户分权管理脚本MySQL用户权限体系配置思路_Sublime中编辑多用户分权管理脚本MySQL用户权限体系配置思路_Sublime中编辑多用户分权管理脚本MySQL用户权限体系配置思路_Sublime中编辑多用户分权管理脚本

    最小权限原则是mysql用户权限配置的核心,确保每个用户仅拥有必要权限以提升安全性与可维护性。1.明确需求:根据用户角色分配如只读、增删改查或结构修改权限;2.创建用户并编写sql脚本进行权限管理,替代手动输入命令,提高效率与一致性;3.使用sublime text等编辑器提升脚本编写效率,利用语法…

    2026年9月21日 用户投稿
    100
  • 音乐文件占用空间太多怎么办_音乐文件占用空间太多如何整理详细指南

    解决音乐文件占空间问题的关键是压缩与整理:先用软件或在线工具降低比特率压缩体积,再按场景分类、利用元数据自动归集,并通过听歌片段和BPM判断保留内容,避免重复与误删。 音乐文件占空间太多,核心解决办法就两条:一是压缩单个文件体积,二是通过有效分类管理提升使用效率。直接删歌不是长久之计,学会整理和优化…

    2026年9月21日
    000

发表回复

登录后才能评论
关注微信