news 2026/10/3 1:27:35

自然语言推断与数据集:基于 d2l-en 仓库的 SNLI 文本对推理实战指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
自然语言推断与数据集:基于 d2l-en 仓库的 SNLI 文本对推理实战指南
  • 文档
  • 教程
  • 人工智能
  • 深度学习
  • 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.

项目地址:https://gitcode.com/gh_mirrors/d2/d2l-en
点击查看免费下载

导读

自然语言推断(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)中找到实现:

  1. d2l.tokenize(d2l/torch.py):按空格把每行切分为词 token(token='word',默认),也可切换为字符级切分(token='char')。这里的前提与假设都先被 tokenize 成 token 列表的列表。
  2. d2l.Vocab(d2l/torch.py):从前提与假设的全部 token 构建词表。注意min_freq=5表示出现频次低于 5 的 token 会被过滤掉;reserved_tokens=['<pad>']保留了填充 token;词表还自动包含'<unk>'(未知 token,索引 0),__getitem__遇到词表外的 token 时会返回unk索引。这也解释了训练集与测试集共享词表的重要性:测试集中的新 token 会落到<unk>上。
  3. 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.vocab

PyTorch 版本:

#@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 微调章节均直接复用该管线。

练习

  1. 机器翻译长期以来基于输出译文与标准译文之间的表层 $n$-gram 匹配来评估。你能设计一种利用自然语言推断来评估机器翻译结果的度量吗?(提示:可把译文作为前提、参考译文作为假设,统计蕴含/矛盾/中立比例的思路值得尝试。)
  2. 如何改变超参数来减小词表大小?(提示:考察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.

项目地址:https://gitcode.com/gh_mirrors/d2/d2l-en
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/3 1:27:31

答辩PPT制作全指南:从结构设计到现场放映的避坑手册

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/10/3 1:26:01

工业级步进电机控制:DRV8818与PIC18F47K40硬实时协同设计

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/10/3 1:25:56

硬盘DMA真相:不是硬盘自动搬运,而是控制器精密调度

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/10/3 1:25:32

嵌入式I2C一主多从总线设计:从物理层到调试实践

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/10/3 1:25:31

VCS增量编译与分离编译:数字IC验证的编译加速实战指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/10/3 1:25:00

Kubernetes持久化存储实战:从PV/PVC到StorageClass与NFS动态供给

1. 为什么Kubernetes需要一套独立的存储抽象 1.1 先聊聊容器世界里的数据到底有多脆弱 熟悉Kubernetes的朋友应该都对这句话不陌生&#xff1a; Pod是"牲畜"而不是"宠物" 。翻译成人话就是&#xff0c;Pod随时可能被销毁、被重建、被调度到另一台节点上…

作者头像 李华