PyTorch 简单易懂的 Embedding 和 EmbeddingBag - 解析与实践

2024-01-08 11:20

本文主要是介绍PyTorch 简单易懂的 Embedding 和 EmbeddingBag - 解析与实践,希望对大家解决编程问题提供一定的参考价值,需要的开发者们随着小编来一起学习吧!

目录

torch.nn子模块Sparse Layers详解

nn.Embedding

用途

主要参数

注意事项

使用示例

从预训练权重创建嵌入

nn.EmbeddingBag

功能和用途

主要参数

使用示例

从预训练权重创建

总结


torch.nn子模块Sparse Layers详解

nn.Embedding

torch.nn.Embedding 是 PyTorch 中一个重要的模块,用于创建一个简单的查找表,它存储固定字典和大小的嵌入(embeddings)。这个模块通常用于存储单词嵌入并使用索引检索它们。接下来,我将详细解释 Embedding 模块的用途、用法、特点以及如何使用它。

用途

  • 单词嵌入:在自然语言处理中,Embedding 模块用于将单词(或其他类型的标记)映射到一个高维空间,其中相似的单词在嵌入空间中彼此靠近。
  • 特征表示:在非自然语言处理任务中,嵌入可以用于任何类型的分类特征的密集表示。

主要参数

  • num_embeddings(int):嵌入字典的大小。
  • embedding_dim(int):每个嵌入向量的大小。
  • padding_idx(int,可选):如果指定,padding_idx 处的嵌入不会在训练中更新。
  • max_norm(float,可选):如果指定,将重新归一化超过此范数的嵌入向量。
  • norm_type(float,可选):用于max_norm选项的p-范数的p值,默认为2。
  • scale_grad_by_freq(bool,可选):如果为True,将按单词在批次中的频率的倒数来缩放梯度。
  • sparse(bool,可选):如果为True,权重矩阵的梯度将是一个稀疏张量。

注意事项

  • 当使用max_norm参数时,Embedding的前向方法会就地修改权重张量。如果需要对Embedding.weight进行梯度计算,则在调用前向方法前,需要在max_norm不为None时克隆它。
  • 仅有少数优化器支持稀疏梯度。

使用示例

import torch
import torch.nn as nn# 创建一个包含10个大小为3的嵌入的Embedding模块
embedding = nn.Embedding(10, 3)# 一个包含4个索引的2个样本的批次
input = torch.LongTensor([[1, 2, 4, 5], [4, 3, 2, 9]])# 通过Embedding模块获取嵌入
output = embedding(input)

此示例创建了一个嵌入字典大小为10、每个嵌入维度为3的 Embedding 模块。然后它接受一个包含索引的输入张量,并返回对应的嵌入向量。

从预训练权重创建嵌入

还可以使用from_pretrained类方法从预先训练的权重创建Embedding实例:

# 预训练的权重
weight = torch.FloatTensor([[1, 2.3, 3], [4, 5.1, 6.3]])# 从预训练权重创建Embedding
embedding = nn.Embedding.from_pretrained(weight)# 获取索引1的嵌入
input = torch.LongTensor([1])
output = embedding(input)

在这个示例中,Embedding 模块是从一个给定的预训练权重张量创建的。这种方法在迁移学习或使用预先训练好的嵌入时非常有用。

nn.EmbeddingBag

torch.nn.EmbeddingBag 是 PyTorch 中一个高效的模块,用于计算“bags”(即序列或集合)的嵌入的总和或平均值,而无需实例化中间的嵌入。这个模块特别适用于处理具有不同长度的序列,如在自然语言处理任务中处理不同长度的句子或文档。下面我将详细介绍 EmbeddingBag 的功能、用法以及特点。

功能和用途

  • 高效计算EmbeddingBag 直接计算整个包的总和或平均值,比逐个嵌入后再求和或取平均更加高效。
  • 支持不同聚合方式:可以选择 "sum", "mean" 或 "max" 模式来聚合每个包中的嵌入。
  • 支持加权聚合EmbeddingBag 还支持为每个样本指定权重,在 "sum" 模式下进行加权求和。

主要参数

  • num_embeddings(int):嵌入字典的大小。
  • embedding_dim(int):每个嵌入向量的大小。
  • max_norm(float,可选):如果给定,将重新规范化超过此范数的嵌入向量。
  • mode(str,可选):聚合模式,可以是 "sum"、"mean" 或 "max"。
  • sparse(bool,可选):如果为True,权重矩阵的梯度将是一个稀疏张量。
  • padding_idx(int,可选):如果指定,padding_idx 处的嵌入将不会在训练中更新。

使用示例

import torch
import torch.nn as nn# 创建一个包含10个大小为3的嵌入的EmbeddingBag模块
embedding_bag = nn.EmbeddingBag(10, 3, mode='mean')# 一个示例包含4个索引的输入
input = torch.tensor([1, 2, 4, 5, 4, 3, 2, 9], dtype=torch.long)# 指定每个包的开始索引
offsets = torch.tensor([0, 4], dtype=torch.long)# 通过EmbeddingBag模块获取嵌入
output = embedding_bag(input, offsets)

在这个示例中,创建了一个嵌入字典大小为10、每个嵌入维度为3的 EmbeddingBag 模块,并设置为 "mean" 模式。输入是一个索引序列,offsets 指定了每个包的开始位置。EmbeddingBag 会计算每个包的平均嵌入向量。

从预训练权重创建

EmbeddingBag 也可以从预训练的权重创建:

# 预训练的权重
weight = torch.FloatTensor([[1, 2.3, 3], [4, 5.1, 6.3]])# 从预训练权重创建EmbeddingBag
embedding_bag = nn.EmbeddingBag.from_pretrained(weight)# 获取索引1的嵌入
input = torch.LongTensor([[1, 0]])
output = embedding_bag(input)

 这种方法在需要使用预先训练好的嵌入或在迁移学习中非常有用。EmbeddingBag 通过高效地处理不同长度的序列数据,在自然语言处理等领域中发挥着重要作用。

总结

 本篇博客探讨了 PyTorch 中的 nn.Embeddingnn.EmbeddingBag 两个关键模块,它们是处理和表示离散数据特征的强大工具。nn.Embedding 提供了一种有效的方式来将单词或其他类型的标记映射到高维空间中,而 nn.EmbeddingBag 以其独特的方式处理变长序列,通过聚合嵌入来提高计算效率。这两个模块不仅在自然语言处理中发挥关键作用,也适用于其他需要稠密特征表示的任务。此外,这些模块支持从预训练权重初始化,使其在迁移学习和复杂模型训练中极为重要。综上所述,nn.Embeddingnn.EmbeddingBag 是理解和应用 PyTorch 中嵌入层的基础。

这篇关于PyTorch 简单易懂的 Embedding 和 EmbeddingBag - 解析与实践的文章就介绍到这儿,希望我们推荐的文章对编程师们有所帮助!



http://www.chinasem.cn/article/583294

相关文章

网页解析 lxml 库--实战

lxml库使用流程 lxml 是 Python 的第三方解析库,完全使用 Python 语言编写,它对 XPath表达式提供了良好的支 持,因此能够了高效地解析 HTML/XML 文档。本节讲解如何通过 lxml 库解析 HTML 文档。 pip install lxml lxm| 库提供了一个 etree 模块,该模块专门用来解析 HTML/XML 文档,下面来介绍一下 lxml 库

基于MySQL Binlog的Elasticsearch数据同步实践

一、为什么要做 随着马蜂窝的逐渐发展,我们的业务数据越来越多,单纯使用 MySQL 已经不能满足我们的数据查询需求,例如对于商品、订单等数据的多维度检索。 使用 Elasticsearch 存储业务数据可以很好的解决我们业务中的搜索需求。而数据进行异构存储后,随之而来的就是数据同步的问题。 二、现有方法及问题 对于数据同步,我们目前的解决方案是建立数据中间表。把需要检索的业务数据,统一放到一张M

csu 1446 Problem J Modified LCS (扩展欧几里得算法的简单应用)

这是一道扩展欧几里得算法的简单应用题,这题是在湖南多校训练赛中队友ac的一道题,在比赛之后请教了队友,然后自己把它a掉 这也是自己独自做扩展欧几里得算法的题目 题意:把题意转变下就变成了:求d1*x - d2*y = f2 - f1的解,很明显用exgcd来解 下面介绍一下exgcd的一些知识点:求ax + by = c的解 一、首先求ax + by = gcd(a,b)的解 这个

hdu2289(简单二分)

虽说是简单二分,但是我还是wa死了  题意:已知圆台的体积,求高度 首先要知道圆台体积怎么求:设上下底的半径分别为r1,r2,高为h,V = PI*(r1*r1+r1*r2+r2*r2)*h/3 然后以h进行二分 代码如下: #include<iostream>#include<algorithm>#include<cstring>#include<stack>#includ

【C++】_list常用方法解析及模拟实现

相信自己的力量,只要对自己始终保持信心,尽自己最大努力去完成任何事,就算事情最终结果是失败了,努力了也不留遗憾。💓💓💓 目录   ✨说在前面 🍋知识点一:什么是list? •🌰1.list的定义 •🌰2.list的基本特性 •🌰3.常用接口介绍 🍋知识点二:list常用接口 •🌰1.默认成员函数 🔥构造函数(⭐) 🔥析构函数 •🌰2.list对象

usaco 1.3 Prime Cryptarithm(简单哈希表暴搜剪枝)

思路: 1. 用一个 hash[ ] 数组存放输入的数字,令 hash[ tmp ]=1 。 2. 一个自定义函数 check( ) ,检查各位是否为输入的数字。 3. 暴搜。第一行数从 100到999,第二行数从 10到99。 4. 剪枝。 代码: /*ID: who jayLANG: C++TASK: crypt1*/#include<stdio.h>bool h

uva 10387 Billiard(简单几何)

题意是一个球从矩形的中点出发,告诉你小球与矩形两条边的碰撞次数与小球回到原点的时间,求小球出发时的角度和小球的速度。 简单的几何问题,小球每与竖边碰撞一次,向右扩展一个相同的矩形;每与横边碰撞一次,向上扩展一个相同的矩形。 可以发现,扩展矩形的路径和在当前矩形中的每一段路径相同,当小球回到出发点时,一条直线的路径刚好经过最后一个扩展矩形的中心点。 最后扩展的路径和横边竖边恰好组成一个直

【生成模型系列(初级)】嵌入(Embedding)方程——自然语言处理的数学灵魂【通俗理解】

【通俗理解】嵌入(Embedding)方程——自然语言处理的数学灵魂 关键词提炼 #嵌入方程 #自然语言处理 #词向量 #机器学习 #神经网络 #向量空间模型 #Siri #Google翻译 #AlexNet 第一节:嵌入方程的类比与核心概念【尽可能通俗】 嵌入方程可以被看作是自然语言处理中的“翻译机”,它将文本中的单词或短语转换成计算机能够理解的数学形式,即向量。 正如翻译机将一种语言

poj 1113 凸包+简单几何计算

题意: 给N个平面上的点,现在要在离点外L米处建城墙,使得城墙把所有点都包含进去且城墙的长度最短。 解析: 韬哥出的某次训练赛上A出的第一道计算几何,算是大水题吧。 用convexhull算法把凸包求出来,然后加加减减就A了。 计算见下图: 好久没玩画图了啊好开心。 代码: #include <iostream>#include <cstdio>#inclu

uva 10130 简单背包

题意: 背包和 代码: #include <iostream>#include <cstdio>#include <cstdlib>#include <algorithm>#include <cstring>#include <cmath>#include <stack>#include <vector>#include <queue>#include <map>