- 文档
- 教程
- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】d2l-en
Interactive deep learning book with multi-framework code, math, and discussions. Adopted at 500 universities from 70 countries including Stanford, MIT, Harvard, and Cambridge.
导读
自然语言推断(Natural Language Inference, NLI)回答的是"一个文本序列(hypothesis,假设)能否从另一个文本序列(premise,前提)中被推断出来"这一逻辑关系判定问题。与情感分析只对单条文本做分类不同,NLI 需要对成对文本序列进行推理,是信息检索、开放域问答等上层应用的基础能力。本文以 d2l-en 仓库中 自然语言推断与数据集章节 为骨架,完整讲解 NLI 的三分类任务定义、斯坦福 SNLI 语料的下载与解析、基于 Gluon 与 PyTorch 的自定义数据集类实现,并结合仓库源码(d2l/torch.py、d2l/mxnet.py)剖析download_extract、Vocab、truncate_pad等底层工具的真实调用链。读完本文,你将掌握一套可直接复用的 SNLI 数据处理管线:从原始 zip 下载、制表符解析、文本清洗,到词表构建、定长截断填充,再到返回DataLoader迭代器与形状正确的成对输入。
自然语言推断:判定一对文本之间的逻辑关系
在 情感分析章节 中,我们讨论了将单条文本序列分类到预定义类别(如情感极性)的任务。但当我们需要判断"一个句子能否从另一个句子推断出来",或通过识别语义等价的句子来消除冗余信息时,仅对单条序列分类是不够的——我们需要能够对成对的文本序列进行推理。
自然语言推断研究的是:假设(hypothesis)能否从前提(premise)中推断出来,其中两者都是文本序列。换句话说,自然语言推断判定一对文本序列之间的逻辑关系。这种关系通常分为三类:
- 蕴含(Entailment):假设可以从前提中推断出来。
- 矛盾(Contradiction):前提中可以推断出假设的否定。
- 中立(Neutral):其余所有情况。
自然语言推断也被称为识别文本蕴含(Recognizing Textual Entailment, RTE)任务。以下面三组示例来说明三类标签:
- 蕴含:由于前提中的 "hugging one another"(互相拥抱)可以推断出假设中的 "showing affection"(表达感情),因此这对句子被标注为entailment。
- 前提:Two women are hugging each other.
- 假设:Two women are showing affection.
- 矛盾:因为 "running the coding example"(运行代码示例)表明 "not sleeping"(没有睡觉),而不是 "sleeping"(睡觉),因此是contradiction。
- 前提:A man is running the coding example from Dive into Deep Learning.
- 假设:The man is sleeping.
- 中立:从 "are performing for us"(为我们表演)这一事实,既无法推断出 "famous"(出名),也无法推断出 "not famous"(不出名),因此是neutrality。
- 前提:The musicians are performing for us.
- 假设:The musicians are famous.
自然语言推断一直是理解自然语言的核心课题,其应用广泛,涵盖从信息检索到开放域问答等多个领域。为了研究该问题,我们首先考察一个流行的自然语言推断基准数据集。
斯坦福自然语言推断(SNLI)语料
斯坦福自然语言推断(SNLI)语料库是包含50 万余条标注英文句子对的集合(引自Bowman.Angeli.Potts.ea.2015)。在 d2l-en 仓库中,我们下载并解压 SNLI 数据集到../data/snli_1.0路径(相对仓库根目录即data/snli_1.0)。
首先,在d2l.DATA_HUB中注册该数据集。仓库源码中DATA_HUB定义于 d2l/torch.py(DATA_HUB = dict()),MXNet 与 PyTorch 两个后端均使用同一份注册信息,并记录了下载地址与 SHA-1 校验值:
#@tab mxnet from d2l import mxnet as d2l from mxnet import gluon, np, npx import os import re npx.set_np() #@save d2l.DATA_HUB['SNLI'] = ( 'https://nlp.stanford.edu/projects/snli/snli_1.0.zip', '9fcde07509c7e87ec61c640c1b2753d9041758e4') data_dir = d2l.download_extract('SNLI')#@tab pytorch from d2l import torch as d2l import torch from torch import nn import os import re #@save d2l.DATA_HUB['SNLI'] = ( 'https://nlp.stanford.edu/projects/snli/snli_1.0.zip', '9fcde07509c7e87ec61c640c1b2753d9041758e4') data_dir = d2l.download_extract('SNLI')源码层面的下载与解压链路
d2l.download_extract的真实实现在 d2l/torch.py 中:它先调用download完成文件下载,再按扩展名(.zip、.tar、.gz)解压到基目录。download函数(d2l/torch.py)的细节值得注意:
- 若传入的
url不以http开头,则视为DATA_HUB中的键,自动取出(url, sha1_hash)元组——这正是download_extract('SNLI')的用法; - 文件保存到
../data/目录下(download的默认folder='../data'); - 若文件已存在且 SHA-1 哈希与注册值一致,则直接命中缓存返回,避免重复下载;
- 否则通过
requests.get(url, stream=True, verify=True)流式下载并写入文件。
也就是说,注册的哈希值9fcde0...是完整性校验的关键,数据损坏或来源不一致时会触发重新下载。
读取数据集:read_snli
原始 SNLI 数据集包含比我们实验所需丰富得多的信息。因此定义一个read_snli函数,只提取其中一部分,然后返回前提列表、假设列表及对应的标签列表:
#@tab all #@save def read_snli(data_dir, is_train): """Read the SNLI dataset into premises, hypotheses, and labels.""" def extract_text(s): # Remove information that will not be used by us s = re.sub('\\(', '', s) s = re.sub('\\)', '', s) # Substitute two or more consecutive whitespace with space s = re.sub('\\s{2,}', ' ', s) return s.strip() label_set = {'entailment': 0, 'contradiction': 1, 'neutral': 2} file_name = os.path.join(data_dir, 'snli_1.0_train.txt' if is_train else 'snli_1.0_test.txt') with open(file_name, 'r') as f: rows = [row.split('\t') for row in f.readlines()[1:]] premises = [extract_text(row[1]) for row in rows if row[0] in label_set] hypotheses = [extract_text(row[2]) for row in rows if row[0] in label_set] labels = [label_set[row[0]] for row in rows if row[0] in label_set] return premises, hypotheses, labels该函数的关键设计点:
- 文本清洗:
extract_text用正则去掉括号字符()(原始 SNLI 文本中常包含括号括起来的标注信息),再把连续两个及以上的空白字符压缩为单个空格,最后strip()去首尾空白; - 标签映射:
label_set将'entailment'、'contradiction'、'neutral'分别映射为整数0、1、2,供后续分类模型直接使用; - 列结构:SNLI 原始
tsv文件以\t分隔,跳过首行表头(f.readlines()[1:]),第 1 列为标签、第 2 列为前提、第 3 列为假设; - 过滤无效行:只有标签落在
label_set中的行才被保留,其余(如标注为'-'的行)被丢弃。
现在打印前 3 对前提和假设及其标签("0"、"1"、"2" 分别对应 "entailment"、"contradiction"、"neutral"):
#@tab all train_data = read_snli(data_dir, is_train=True) for x0, x1, y in zip(train_data[0][:3], train_data[1][:3], train_data[2][:3]): print('premise:', x0) print('hypothesis:', x1) print('label:', y)训练集约含55 万对,测试集约含1 万对。下面的统计显示,三类标签在训练集与测试集中都是均衡的:
#@tab all test_data = read_snli(data_dir, is_train=False) for data in [train_data, test_data]: print([[row for row in data[2]].count(i) for i in range(3)])均衡的类别分布意味着可以直接用准确率评估模型,而不必担心类别先验偏差。
定义加载数据集的类:SNLIDataset
下面通过继承 Gluon 的Dataset类来定义一个加载 SNLI 数据集的类。构造函数中的num_steps参数指定文本序列的长度,使每个小批量的序列形状一致。换言之,较长的序列中第num_steps个 token 之后的部分会被截断,而较短的序列则会追加特殊 token"<pad>"直到长度达到num_steps。通过实现__getitem__函数,可以用索引idx任意访问前提、假设和标签。
MXNet 实现(继承gluon.data.Dataset):
#@tab mxnet #@save class SNLIDataset(gluon.data.Dataset): """A customized dataset to load the SNLI dataset.""" def __init__(self, dataset, num_steps, vocab=None): self.num_steps = num_steps all_premise_tokens = d2l.tokenize(dataset[0]) all_hypothesis_tokens = d2l.tokenize(dataset[1]) if vocab is None: self.vocab = d2l.Vocab(all_premise_tokens + all_hypothesis_tokens, min_freq=5, reserved_tokens=['<pad>']) else: self.vocab = vocab self.premises = self._pad(all_premise_tokens) self.hypotheses = self._pad(all_hypothesis_tokens) self.labels = np.array(dataset[2]) print('read ' + str(len(self.premises)) + ' examples') def _pad(self, lines): return np.array([d2l.truncate_pad( self.vocab[line], self.num_steps, self.vocab['<pad>']) for line in lines]) def __getitem__(self, idx): return (self.premises[idx], self.hypotheses[idx]), self.labels[idx] def __len__(self): return len(self.premises)PyTorch 实现(继承torch.utils.data.Dataset):
#@tab pytorch #@save class SNLIDataset(torch.utils.data.Dataset): """A customized dataset to load the SNLI dataset.""" def __init__(self, dataset, num_steps, vocab=None): self.num_steps = num_steps all_premise_tokens = d2l.tokenize(dataset[0]) all_hypothesis_tokens = d2l.tokenize(dataset[1]) if vocab is None: self.vocab = d2l.Vocab(all_premise_tokens + all_hypothesis_tokens, min_freq=5, reserved_tokens=['<pad>']) else: self.vocab = vocab self.premises = self._pad(all_premise_tokens) self.hypotheses = self._pad(all_hypothesis_tokens) self.labels = torch.tensor(dataset[2]) print('read ' + str(len(self.premises)) + ' examples') def _pad(self, lines): return torch.tensor([d2l.truncate_pad( self.vocab[line], self.num_steps, self.vocab['<pad>']) for line in lines]) def __getitem__(self, idx): return (self.premises[idx], self.hypotheses[idx]), self.labels[idx] def __len__(self): return len(self.premises)底层工具函数解析
该类的三个核心依赖均可在仓库 d2l/torch.py(MXNet 版对应 d2l/mxnet.py)中找到实现:
d2l.tokenize(d2l/torch.py):按空格把每行切分为词 token(token='word',默认),也可切换为字符级切分(token='char')。这里的前提与假设都先被 tokenize 成 token 列表的列表。d2l.Vocab(d2l/torch.py):从前提与假设的全部 token 构建词表。注意min_freq=5表示出现频次低于 5 的 token 会被过滤掉;reserved_tokens=['<pad>']保留了填充 token;词表还自动包含'<unk>'(未知 token,索引 0),__getitem__遇到词表外的 token 时会返回unk索引。这也解释了训练集与测试集共享词表的重要性:测试集中的新 token 会落到<unk>上。d2l.truncate_pad(d2l/torch.py):当len(line) > num_steps时截断到前num_steps个 token;否则在末尾用padding_token(此处为vocab['<pad>']的索引)补齐到num_steps长度。_pad方法正是对每个 token 序列应用vocab[line](token 转索引)后再做截断/填充,最终得到形状为(样本数, num_steps)的整数张量。
此外,__getitem__返回的元组结构(premises[idx], hypotheses[idx]), labels[idx]是 NLI 任务区别于情感分析的关键:每个样本包含两个输入(前提与假设)和一个标签。
整合全部流程:load_data_snli
现在调用read_snli函数和SNLIDataset类来下载 SNLI 数据集,并返回训练集和测试集的DataLoader实例,以及训练集的词表。值得强调的是,必须使用从训练集构建的词表作为测试集的词表。这样一来,测试集中出现的任何新 token 对在训练集上训练的模型而言都是未知的(被映射为<unk>),从而保证评估的真实性,避免数据泄露。
MXNet 版本:
#@tab mxnet #@save def load_data_snli(batch_size, num_steps=50): """Download the SNLI dataset and return data iterators and vocabulary.""" num_workers = d2l.get_dataloader_workers() data_dir = d2l.download_extract('SNLI') train_data = read_snli(data_dir, True) test_data = read_snli(data_dir, False) train_set = SNLIDataset(train_data, num_steps) test_set = SNLIDataset(test_data, num_steps, train_set.vocab) train_iter = gluon.data.DataLoader(train_set, batch_size, shuffle=True, num_workers=num_workers) test_iter = gluon.data.DataLoader(test_set, batch_size, shuffle=False, num_workers=num_workers) return train_iter, test_iter, train_set.vocabPyTorch 版本:
#@tab pytorch #@save def load_data_snli(batch_size, num_steps=50): """Download the SNLI dataset and return data iterators and vocabulary.""" num_workers = d2l.get_dataloader_workers() data_dir = d2l.download_extract('SNLI') train_data = read_snli(data_dir, True) test_data = read_snli(data_dir, False) train_set = SNLIDataset(train_data, num_steps) test_set = SNLIDataset(test_data, num_steps, train_set.vocab) train_iter = torch.utils.data.DataLoader(train_set, batch_size, shuffle=True, num_workers=num_workers) test_iter = torch.utils.data.DataLoader(test_set, batch_size, shuffle=False, num_workers=num_workers) return train_iter, test_iter, train_set.vocab参数与行为说明:
num_workers:来自d2l.get_dataloader_workers(),仓库实现固定返回 4 个进程读取数据(d2l/torch.py),用于加速数据装载;batch_size与num_steps=50:num_steps是默认的序列长度上限,可通过调用时传参覆盖;- 打乱策略:训练集
shuffle=True打乱顺序,测试集shuffle=False保持顺序以便稳定评估。
这里将批量大小设为 128、序列长度设为 50,调用load_data_snli获取数据迭代器和词表,然后打印词表大小:
#@tab all train_iter, test_iter, vocab = load_data_snli(128, 50) len(vocab)接着打印第一个小批量的形状。与情感分析不同,这里有两个输入X[0]和X[1],分别代表前提对与假设对:
#@tab all for X, Y in train_iter: print(X[0].shape) print(X[1].shape) print(Y.shape) break在batch_size=128、num_steps=50下,输出形状应为X[0]: (128, 50)、X[1]: (128, 50)、Y: (128,)——前提与假设各是一个定长 50 的索引序列批量,标签是长度为 128 的整数向量。
在仓库中的实际消费场景
该数据集加载管线并非孤立存在,而是后续两个 NLI 模型章节的直接数据来源,这印证了其接口设计的通用性:
- 自然语言推断:使用注意力 在第 335 行直接调用
d2l.load_data_snli(batch_size, num_steps),把返回的迭代器与词表喂给基于注意力与 MLP 的可分解注意力模型; - 自然语言推断:微调 BERT 在第 262-277 行复用
d2l.read_snli(data_dir, True/False)读取原始三元组,再按 BERT 的输入格式(拼接两个序列)重新封装数据集。
从源码结构看,read_snli之所以被设计成"只返回前提、假设、标签三个列表"的纯函数,正是为了同时服务"定长截断填充"(本节的SNLIDataset)与"BERT 拼接输入"(后续章节)两种不同的封装方式。这也为读者自己扩展新的 NLI 数据封装(例如加入注意力掩码、变长序列打包)提供了清晰的切入点。
小结
- 自然语言推断研究假设能否从前提中推断出来,两者均为文本序列;
- 在自然语言推断中,前提与假设之间的关系包括蕴含(entailment)、矛盾(contradiction)和中立(neutral)三类;
- 斯坦福自然语言推断(SNLI)语料库是自然语言推断的流行基准数据集;
- 仓库提供了完整的 SNLI 处理管线:
read_snli(读取+清洗)→SNLIDataset(分词、建词表、截断填充)→load_data_snli(返回训练/测试迭代器与词表),后续注意力模型与 BERT 微调章节均直接复用该管线。
练习
- 机器翻译长期以来基于输出译文与标准译文之间的表层 $n$-gram 匹配来评估。你能设计一种利用自然语言推断来评估机器翻译结果的度量吗?(提示:可把译文作为前提、参考译文作为假设,统计蕴含/矛盾/中立比例的思路值得尝试。)
- 如何改变超参数来减小词表大小?(提示:考察
Vocab的min_freq参数、num_steps序列长度,以及分词粒度 word/char 对词表规模的影响。)
- 文档
- 教程
- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】d2l-en
Interactive deep learning book with multi-framework code, math, and discussions. Adopted at 500 universities from 70 countries including Stanford, MIT, Harvard, and Cambridge.
相关推荐
深入理解自然语言推断与SNLI数据集
深入理解自然语言推断与SNLI数据集 自然语言处理 NLP 领域中,自然语言推断 Natural Language Inference, NLI 是一项基础且重
人工智能深度学习机器学习教程深入理解自然语言推理与SNLI数据集
深入理解自然语言推理与SNLI数据集 自然语言推理 Natural Language Inference, NLI 是自然语言处理领域中的一个重要任务,它研究如
文档教程人工智能深度学习NLP计算机视觉强化学习自然语言推断与SNLI数据集实战:《动手学深度学习》文本对分类的数据准备全解析
自然语言推断与SNLI数据集实战:《动手学深度学习》文本对分类的数据准备全解析 自然语言推断(Natural Language Inference, NLI)是
人工智能深度学习机器学习教程
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考