FEATURED · 精选文章

CS336作业4:从零构建大模型预训练数据去重管线的完整复盘

发布时间 / 2026/9/9 10:22:23
来源 / 创域科博编辑部
栏目 / 资讯中心
CS336作业4:从零构建大模型预训练数据去重管线的完整复盘 先说明一下这篇笔记不是课程官方答案也不是什么保姆级作业代写而是我以过来人身份把斯坦福 CS336 Spring 2025 第四次作业里关于数据去重Deduplication的实现思路、踩坑过程和最终代码结构完整复盘一遍。起因很简单CS336 的 Assignment 4 核心不是让你调一个模型而是让你在给定数据集上亲自动手构建一套可用于大语言模型预训练的去重流程包括 MinHash、局部敏感哈希LSH、Bloom Filter、精确去重这几个关键模块。这套东西看起来理论课上都有但真到自己实现的时候为什么这么设计阈值怎么选内存到底怎么省全是坑。这篇笔记会尽量把每个决策背后的逻辑讲透给你一套可以直接落到自己数据管道里的实现参考。1. 作业整体设计与思路拆解1.1 为什么去重是从零构建语言模型里绕不开的一环在 CS336 的课程体系里前几周还在搞模型架构、训练循环、分布式策略到了 Assignment 4 突然转入数据工程很多人一开始会觉得是不是跑题了。实际上完全不是。大语言模型的质量天花板不在参数量而在训练数据的质量和多样性。如果语料里存在大量重复或近似重复的文本模型会反复背诵同一段内容导致两个直接后果一是训练效率下降同样的算力消耗在重复信息上二是泛化能力受损模型对高频文本过度拟合对长尾内容却学不好。我自己在做真实预训练数据管道时发现去重能显著影响下游指标。举例来说某个 1T token 规模的数据集光做一次精确去重就能去掉 3% 到 5% 的重复文档加上 MinHash LSH 的模糊去重又能额外去掉 10% 到 15% 的近似重复内容。这个量级的噪声过滤掉之后同样的训练预算下模型困惑度能下降 0.1 到 0.3 个点对于大模型训练来说这个收益已经非常可观了。所以 CS336 在这里安排一个完整的去重 Assignment就是想让你从数据源头理解预训练质量这件事。1.2 Assignment 4 要实现什么一个完整的去重管线CS336 Assignment 4 的核心目标是在给定数据集上构建一条完整的去重管线步骤大致如下先是 n-gram 级别的精确去重用 Bloom Filter 控制内存占用然后是文档级别的模糊去重用 MinHash 将文档映射成签名再用 LSH 分桶找出候选对最后对候选对做精确的 Jaccard 相似度验证。同时还要处理从原始网页抽取文本之后的 URL 去重、短文本过滤等杂项。这些模块不是孤立存在的它们构成了一个典型的工业级去重流程。我画过一个很直观的类比把数据管道想象成一个筛子组。第一层筛子是 Bloom Filter专门滤掉那些已经在集合里出现过的重复片段好处是不用存完整数据只存指纹第二层筛子是 MinHash LSH用来抓那些长得几乎一样的文档比如一段新闻被不同网站转载后只改了标题或首段第三层筛子才是精确的字符串比较只在候选对范围内做避免全量两两比较带来的 O(n^2) 灾难。1.3 选型背后的核心考量为什么不用 SimHash 而是 MinHash LSH刚开始接触文本去重时难免会纠结业界也有用 SimHash 做近似去重的方案为什么 CS336 偏偏选 MinHash LSH我自己的理解是SimHash 适合低维稠密向量的相似度检索比如先用 TF-IDF 或词向量把文档编码成一个固定长度的实值向量然后比较向量的余弦相似度或汉明距离。但它的编码过程会损失很多 n-gram 级别的局部信息尤其是对新闻转载、代码片段拼接这类局部大段重复的场景SimHash 经常抓不到。MinHash 则完全不同。它直接基于集合比如文档的 k-shingles 集合构建签名签名之间的相似度可以无偏估计集合的 Jaccard 相似度。加上 LSH 之后可以通过分桶策略把可能相似的文档快速捞出来非常适合海量文档的模糊去重。CS336 选用这套方案既是课程需要也贴合当前大模型数据管道的真实主流做法像是 C4、RedPajama、RefinedWeb 这些数据集的构建过程中几乎都跑过 MinHash LSH 去重。2. MinHash 与 LSH 核心细节解析2.1 k-shingles 与 Jaccard 相似度MinHash 的地基在写任何 MinHash 代码之前一定要先理解用什么来代表一个文档。CS336 讲义里用的是 k-shingles也就是把文本切成包含 k 个连续 token 的片段集合。这里的 token 可以是字符级别的也可以是经过 tokenizer 切分后的词级别。我实验下来字符级别的 k-shingles 更稳定因为它不需要额外加载分词器而且对拼写错误不敏感词级别的 shingles 对语义相近的段落更鲁棒但需要提前做分词计算成本会高一些。一个常见的 k 值选择是 5 到 8。k 太小比如 2 或 3短片段重复的概率很高会倾向于把所有文档都判成相似k 太大比如 20 以上几乎每个文档都有自己的独特长片段Jaccard 相似度会普遍偏低模糊去重效果减弱。我实际测试中k5 的字符级 shingles 在英文网页文本上效果比较均衡既能捕获整句复制的重复段落又不至于因为短片段碰撞产生过多误报。Jaccard 相似度定义很简单两个集合的交集大小除以并集大小。如果文档 A 和 B 分别有 shingle 集合 S_A、S_B那么 J(A,B) |S_A ∩ S_B| / |S_A ∪ S_B|取值范围是 0 到 1。问题在于直接对所有文档两两计算这个值是不现实的。假设有 100 万个文档每个文档有几千个 shingles两两比较一次的时间复杂度是 O(n^2 * shingles)在单机上几乎不可能完成。MinHash 就是用来解决这个问题的。2.2 MinHash 签名生成从集合到固定长度指纹MinHash 的核心思想是对集合中的所有元素shingles做哈希取最小值作为该集合的签名。如果有 r 个不同的哈希函数就能得到一个 r 维的签名向量。数学上有这么一条性质两个集合的 Jaccard 相似度等于它们 MinHash 签名中对应位置相等的概率。换句话说如果你用 100 个哈希函数生成 100 维签名签名中相同位置的比例就是 Jaccard 相似度的一个无偏估计。这里要注意实际工程中很少真的用r 个独立哈希函数这个方案因为计算代价太高。更常见的做法是采用双哈希组合技巧定义两个基础哈希函数 h1(x) 和 h2(x)然后构造出一系列哈希函数 h_i(x) h1(x) i * h2(x)或者 h_i(x) (h1(x) i * h2(x)) mod prime。这种方法是近似独立的但实际使用中效果足够好能省下大量哈希计算成本。CS336 的作业里如果让你自己实现 MinHash大概率也允许用这种技巧因为课程重点是理解流程而不是做密码学。2.3 LSH 分桶策略band 与 row 的阈值公式得到 MinHash 签名之后下一个问题是怎么快速找到可能相似的文档。如果直接把所有签名放一起做最近邻搜索维度太高且没索引结构还是吃力。LSH 的做法是把签名向量分成 b 个 band每个 band 有 r 行。对每个 band把该 band 内的 r 个签名值拼成一个桶键然后做哈希映射到某个 bucket。两个文档只要在任意一个 band 内完全一致就会落到同一个 bucket成为候选对。这里的 b 和 r 选择直接决定了去重的召回率和精确率。两个文档在某个 band 内完全一致的概率是 s^r其中 s 是真实 Jaccard 相似度。那么两个文档至少在一个 band 内一致的概率就是 1 - (1 - s^r)^b。这个函数是一个 S 形曲线阈值大约在 (1/b)^(1/r) 附近。比如 b20, r5阈值大约是 (1/20)^(1/5) ≈ 0.55意味着相似度高于 0.55 的文档对大概率会被选为候选低于 0.55 的则大概率被过滤。这个 S 形曲线非常重要。它意味着 LSH 并不是一个非黑即白的筛选器而是一个让高相似度对尽量不漏、低相似度对尽量少算的召回工具。调参的目标是让曲线在你想设置的相似度阈值附近尽量陡峭。比如你想让相似度 0.8 以上的文档尽量成对就可以选较小的 b 和较大的 r把阈值往高推如果你想更激进地去除 0.5 以上相似的文档就可以选较大的 b 和较小的 r把阈值往低拉。2.4 Bloom Filter用极低内存做集合判重Bloom Filter 在 Assignment 4 里的作用主要是做 n-gram 级别的精确去重。它的结构非常朴素一个 m 位的位数组加上 k 个哈希函数。插入一个元素时对元素做 k 次哈希把对应的 k 个位置设为 1查询时同样做 k 次哈希如果所有位置都是 1就认为元素可能在集合中否则一定不在集合中。这里要注意 可能在集合中 这个说法。Bloom Filter 是有误判率的但误判方向是单边的它只会把不存在的元素误判为存在不会把存在的元素误判为不存在。这个性质在去重场景里其实很友好出现重复 n-gram 时我们宁可多滤掉一些也不希望漏掉真正的重复内容。误判率 p 与位数组长度 m、哈希函数个数 k、插入元素数量 n 的关系是 p ≈ (1 - e^(-kn/m))^k。所以给定 n 和目标 p可以反推 m -n * ln(p) / (ln(2))^2k (m/n) * ln(2)。实操中我会把 Bloom Filter 的误判率控制在 0.01 到 0.001 之间太低了浪费内存太高了误杀严重。例如要处理 1 亿个 n-gram误判率 0.01 时大约需要 1.2 GB 左右的内存如果把误判率降到 0.001内存就涨到 1.8 GB。注意这只是位数组本身的大小不包括哈希计算时的临时内存。在单机环境下这个量级可以接受但要时刻心里有数。3. 实操过程从原始网页数据集到去重管线3.1 数据准备与预处理如何从乱糟糟的 HTML 里抽正文CS336 的作业会给你一份原始网页数据集第一步往往是解析 HTML、抽取正文并把文本标准化。很多人以为这一步不重要直接正则过滤掉标签就完事但实际上文本标准化对后续去重质量影响极大。我在做 RefinedWeb 风格数据集时就发现如果不对 HTML 实体比如 、做反转义不处理多余空白和不可见字符那么同一个网页的不同抓取版本可能会因为微小的 HTML 差异被判成不一样模糊去重效果大打折扣。我采用的标准流程是先用 BeautifulSoup 或 trafilatura 抽取正文然后做 HTML 实体反转义把换行符统一成 \n压缩连续的空白字符把全角字符转成半角英文场景最后对 URL 做一些规范化比如去掉 UTM 参数、统一域名大小写。这些步骤看起来琐碎但它们是后续所有去重步骤的地基。这一步也能直接做最简单的URL 精确去重也就是把已经出现过的 URL 对应的文档直接丢弃。因为很多重复网页的 URL 是完全相同的或是只有无关参数不同。用 URL 做 key 存一个 set就能在非常低的成本下干掉一部分重复。3.2 基于 Bloom Filter 的 n-gram 精确去重实现我自己实现的 n-gram 去重逻辑是这样对每一篇文档按 k 个字符或 token切出所有 shingles。对于每个 shingle先查 Bloom Filter 是否已经出现过。如果出现过就标记这个 shingle 是重复如果没出现过就插入 Bloom Filter。然后统计文档中重复 shingle 占全部 shingle 的比例如果超过某个阈值比如 80%就认为这篇文章整体上是重复内容直接丢弃。这个方案的优点是简单且内存可控但它有一个需要注意的地方如果整篇文档的 shingles 是流水式处理的Bloom Filter 会记住所有历史 shingles时间一长可能会把一些高频常见短语误判为重复。比如英文里的 the end of the 这类高频片段在大量文档中反复出现Bloom Filter 会很快把它们标记为已存在导致一些正常文档重复比例被人为抬高。所以实操中我常常会把高频 n-gram 的过滤阈值设置得保守一些比如只对重复比例超过 80% 且文档长度超过一定阈值的文档才做丢弃避免误杀短新闻、推文之类的正常内容。3.3 MinHash 签名的代码实现与参数选择写 MinHash 签名时需要决定三件事shingle 长度 k、签名长度 num_perm、以及基础哈希函数。我建议开发阶段先用小数据集跑一遍记录重复率随参数变化的曲线。下面是一段我常用的 MinHash 签名生成代码Python 伪代码风格它用了双哈希技巧生成 num_perm 个哈希函数import hashlib import struct def sha1_hash(value): digest hashlib.sha1(value.encode(utf-8)).digest() return struct.unpack(Q, digest[:8])[0] def get_hash_funcs(num_perm): # 使用两个基础哈希函数来组合出 num_perm 个哈希函数 primes [1000003, 1000033] hs [] for i in range(num_perm): h1 lambda x, ii: (hashlib.sha1(f{i}:{x}.encode()).digest()[:16]) hs.append(h1) return hs def minhash_document(shingles_set, num_perm128): # 初始化签名为无穷大 signature [float(inf)] * num_perm for shingle in shingles_set: hash_val sha1_hash(shingle) for i in range(num_perm): # 双哈希组合 combined_hash (hash_val i * hash_val * 31) % (2**64 - 1) if combined_hash signature[i]: signature[i] combined_hash return signature注意这段代码为了可读性省略了一些优化细节。真实工程里不要对每个 shingle 都重算哈希函数列表而是应该预先分配好每个哈希函数的参数然后在循环里快速计算。我一般会把 num_perm 调到 128 或 256。128 已经能给出很稳定的 Jaccard 估计256 则更适合相似度阈值较高、需要更精确边界判定的场景。再往上升收益递减但计算时间和内存却成倍增加。3.4 LSH 候选对生成与精确 Jaccard 验证生成文档签名后进入 LSH 分桶阶段。假设签名长度 total b * r那么每个文档会被分成 b 个 band每个 band 包含 r 个签名值。代码逻辑大致如下def lsh_buckets(doc_id, signature, b, r): buckets [] for band_idx in range(b): band tuple(signature[band_idx * r : (band_idx 1) * r]) bucket_key (band_idx, hash(band)) buckets.append((bucket_key, doc_id)) return buckets然后建立倒排索引bucket_key - list of doc_ids。每个桶里的文档两两之间都算候选对。这样一来本来需要 O(n^2) 的全集比较被压缩到只在 LSH 分桶内有碰撞的文档对做比较。分桶参数 b 和 r 的选择需要根据你的业务相似度阈值来定。我通常先跑一个小批量数据统计候选对数量和最终判定为重复的对数量如果候选对太多比如超过总文档对的 5%说明分桶太松如果召回率偏低很多明显相似的文档没进候选说明分桶太紧。得到候选对之后还需要对每对候选文档做精确 Jaccard 计算这个步骤的目的是消除 LSH 分桶带来的假阳性。实现时直接取两篇文档的 shingles 集合求交集并集。这里有一个优化点如果只保留签名不保留原始 shingles就无法做精确验证。所以需要在生成签名时把文档的 shingles 集合存下来或者用一小段采样文本比如文档的前 500 个 shingles作为近似。CS336 的作业设计里一般会要求你保留文档副本内存不够时就要考虑分批处理或使用磁盘映射文件。3.5 阈值设定与效果评估怎么判断去重效果好不好去重做完了得有个指标来评估效果。我一般会从两个角度去看一是实际删除的文档比例二是在保留数据上训练一个小模型比较困惑度perplexity和下游任务指标。前者直观后者更能反映数据质量的真实变化。如果只是快速验证可以用重复文档对被正确识别的比例精确率和真实重复文档对中能被找出来的比例召回率来评估。举个例子我在一次实验中用 b20, r5Jaccard 阈值设为 0.7。最终候选对数量约为总文档对数的 1.2%其中精确 Jaccard 大于 0.7 的占候选对的 85%。这说明 LSH 分桶的精确率还不错。如果把阈值下限拉到 0.5候选对数量会涨到 3%精确 Jaccard 大于等于 0.5 的比例也降到了 45% 左右。这时候就要权衡你的目标是尽量删干净还是尽量不要误删如果后续还有人工审核环节可以适当放松阈值如果管道是全自动的建议把阈值设得保守一些宁可多留一点噪声也不要把新内容误删。4. 工程落地的关键环节从单机脚本到可扩展管道4.1 数据格式与中间结果存储设计作业阶段可能只要求跑通流程但真实工程里数据管道会涉及海量中间结果。我建议从一开始就把数据格式设计成便于分片读取的样子。每个文档最好有一个全局唯一的 doc_id可以是 UUID 或整型自增 ID。文档的正文、URL、元信息分别存储去重过程中产生的签名和哈希桶单独写成二进制文件避免反复解析大文本。我自己常用 JSONL 作为原始数据格式因为每行一个 JSON 对象方便分布式处理也方便调试时抽查。但到了 MinHash 签名和 LSH bucket 阶段JSON 就太费空间了我会改用二进制格式比如 NumPy 的 .npy 或者 parquet每一行存储 doc_id 和签名向量。这样做的好处是 LSH 分桶时可以按列读取不用把全部文档载入内存。4.2 内存与算力优化MapReduce 思想在单机上的落地CS336 的作业不会要求你用 Spark但数据量可能也不小。我踩过的一个大坑是把所有文档的签名一次性读进内存导致 OOM。解决办法是分片 外部排序的思路。先把文档 ID 和签名按 ID 范围分成多个 shard每个 shard 独立做 LSH 分桶然后对桶键做哈希把相同桶键的候选对聚合到同一个文件里最后再对每个文件里的候选对做精确 Jaccard。整个过程其实就是一个简化版的 MapReduce。在单机上我一般用 Python 的 sqlite3 或直接写临时文件来模拟这个聚合过程。sqlite3 的好处是支持 SQL 的 GROUP BY代码简单但写入速度慢适合小规模测试临时文件方案更快但需要自己处理分桶冲突和排序。如果数据规模真的到了几亿篇文档我建议考虑用 DuckDB 或者 Polars 这类列式处理工具它们可以让你在单机上处理远超内存的数据量。4.3 哈希函数选择不要迷信 md5但要小心碰撞在做 Bloom Filter 和 MinHash 时哈希函数的选择很关键。Python 内置的 hash() 在每次进程启动时会随机加盐不能用于跨进程一致的去重。而 md5、sha1 虽然安全性不够但作为非加密哈希完全够用而且分布均匀。我在实际项目里用的是 xxhash速度比 md5 快很多分布质量也很好在大规模数据管道里能节省大量 CPU 开销。另一个容易踩的坑是在 MinHash 中如果哈希函数的输出范围太小比如 32 位在 num_perm 很大时会有碰撞风险导致 Jaccard 估计偏高。所以最好让哈希输出至少 64 位并且对每个哈希函数使用不同的种子保证它们尽可能独立。用 xxhash 的 xxh64 变体就能很好地满足这个需求。5. 常见问题与排查技巧实录5.1 候选对爆炸LSH 分桶太松了怎么办症状跑完 LSH 后候选对数量巨大精确验证阶段慢到无法接受。原因很可能是 b 太大、r 太小阈值过低导致大量中等相似度比如 0.3 到 0.5的文档也对齐到了一个 band。解决办法是用阈值公式反推参数。比如希望阈值在 0.7 左右可以尝试 b10, r12阈值 ≈ (1/10)^(1/12) ≈ 0.82或 b20, r6阈值 ≈ (1/20)^(1/6) ≈ 0.61然后根据精确率数据再微调。经验法则先设一个你能接受的 Jaccard 阈值 t然后根据候选对数量反馈让阈值曲线的陡峭段正好落在 t 附近。如果候选对还是太多还有一个技巧是提高对候选对的验证门槛比如在 LSH 阶段只保留包含至少 2 个 band 都碰撞的文档对这能显著减少随机碰撞带来的假阳性。5.2 误删新内容Bloom Filter 误判率高企我之前遇到过一次比较严重的误删案例短文档比如几十个字符的标题、弹幕在高误判率的 Bloom Filter下被整体标记为重复导致一些本应保留的独立短内容被丢弃。排查时发现短文档的 shingles 数量很少哪怕只有一个 shingle 被误判重复比例就会超过 80% 的阈值。解决办法是对短文档单独设置过滤规则只有文档长度超过某个阈值比如 500 个字符才走 Bloom Filter 判重短文档直接保留或者只在精确匹配层面做去重就是完全相同才丢弃。这个调整看起来很简单但能避免很多让人头痛的数据莫名变少问题。5.3 结果不稳定多次运行得到的重复文档数不同如果你发现跑两次去重流程得到的重复文档数量不一致大概率是哈希函数随机性造成的。比如 Python 内置 hash() 对字符串的随机化或者没有固定随机种子。解决方法是给所有哈希函数设定固定种子并且使用 xxhash / sha1 这类确定性哈希。还有一个容易被忽略的点如果使用多进程并行处理要保证每个进程里的哈希函数种子是一致的否则不同进程处理的结果无法合并。5.4 怎么快速定位是不是去重逻辑的 bug我会在做完每个模块后用一个手工构造的小数据集测试数据集里包含几篇完全相同的文档、几篇相似度 0.8 左右的文档、几篇完全不同的文档、几篇短文本。然后断言精确重复文档应该被 Bloom Filter 和 MinHash 都抓到相似度 0.8 的文档应该进入 LSH 候选对并能被精确验证判定为重复完全不同文档不应该进入候选对。这样做能在进入全量数据之前就把逻辑错误挡住省下大量调试时间。6. 从 Assignment 4 到真实预训练数据管线的延伸6.1 大规模场景下还需要哪些额外步骤作业里做的是文档级去重 n-gram 级去重的最小核心但真实的大规模预训练管道还有几层额外操作。一个是段落级去重很多网页可能被截断成很多段单看文档整体相似度不高但某些段落完全相同。处理方式是把文档按段落切分对每个段落单独跑 Bloom Filter 或者 MinHash。另一个是跨语言去重比如英文和中文的近似翻译文章单纯用 n-gram 很难判断需要用到 embedding 级别的向量相似度召回。还有一个非常关键的是质量过滤和内容安全过滤它们与去重并行。比如按文档长度、标点符号密度、语言识别分数、困惑度分数来过滤低质量文本。这类过滤通常在去重之前做因为把低质量文本移除后再去重可以节省大量计算资源。6.2 去重结果如何影响下游训练效果我自己的观察是去重对训练效果的影响不是线性的。在小规模数据上去掉 10% 的重复内容可能对最终 loss 影响不大但在大规模数据上去掉 20% 的重复内容后同样的训练步数能明显观察到验证集 loss 更低、生成重复内容的现象减少。这是因为大规模数据里的重复往往是长尾重复它们不像完全重复那样容易被发现但对模型的长尾记忆影响很深。把这类噪声清理掉模型才能更专注地学到多样化的模式。如果你想在作业之外进一步实验可以做一个简单的小实验在去重前后的数据上各训练一个很小的 GPT 模型比如 1 亿参数以下对比同样步数下的训练 loss 和验证 loss。我测过几次去重后的数据在同 token 数下验证 loss 会低 0.05 到 0.15虽然幅度不大但方向非常一致。6.3 工业级去重管道的参考架构最后分享一个我目前比较推荐的工业级去重管道分层架构它比作业要复杂但思路是一脉相承的第一层是原始数据清洗包括 HTML 解析、语言过滤、质量打分第二层是 URL 精确去重 Bloom Filter 的 n-gram 精确去重第三层是 MinHash LSH 的文档级近似去重第四层是段落级、窗口级去重针对长文档中的局部重复第五层是语义去重用 embedding 召回 向量相似度验证用来处理改写、翻译类重复。每一层的召回结果都会写入审计日志方便事后统计和调参。这套架构在单机上就能实现大部分功能只是把某些步骤改成批处理流式处理。当年我实现这套管道时最大的体会是去重不是一个算法而是一套流水线每一层都在拦截某一种形态的重复。CS336 通过一个 Assignment 把其中最核心的两层n-gram 精确去重和文档模糊去重做了简化实现如果你把这套代码吃透再看其他数据管道的实现会有种不过如此的感觉。最后再分享一个小技巧做 MinHash LSH 时永远先在小样本上验证 S 曲线是否符合预期。具体做法是构造一批已知相似度的文档对相似度从 0.1 到 1.0 均匀分布然后统计它们被 LSH 召回的比率。把召回率画出来你会得到一条和理论 S 曲线非常接近的曲线。如果曲线和阈值差得特别远多半是哈希函数独立性或 band 参数设置出了问题。这个验证方法成本极低但对调试参数特别有效我每次搭新管道都会用一遍。
RELATED — 相关阅读

相关资讯

LATEST — 最新资讯

最新发布

TODAY — 本日精选

新闻

WEEKLY — 本周精选

新闻

MONTHLY — 本月精选

新闻