Datasets 批量 map 怎么把一行拆成多行?避免列长不一致并保留来源

用 batched map 将文本拆成多行,同步复制来源与标签,移除旧列并排查 ArrowInvalid 列长不一致。

Datasets 的 batched map 可以让一批 N 条输入产生 M 条输出。把一篇文本拆成多个训练片段时,关键是让输出的每一列长度相同,同时移除仍保持旧行数的列,并为每个片段复制来源 ID 和适用标签。

为什么会出现 ArrowInvalid

官方 Batch mapping 文档明确允许输入输出行数不同,但同一个输出批次里的所有列必须等长。比如原表有 2 条 id,新函数只返回 5 条 chunk,未移除的旧 id 仍是 2 条,表就无法组成矩形,可能报“expected length 2 but got length 5”。

Datasets 批量 map 怎么把一行拆成多行?避免列长不一致并保留来源

示例于 2026-10-01 在 Windows、Python 3.11.15、Datasets 4.8.3 环境实际执行通过;数据全部人为构造,只验证数据处理接口,没有训练模型或测量业务效果。安装可参考 Datasets 官方安装说明。先在新的练习目录建立环境:

python -m venv .venv
# Windows 使用 .venv\Scripts\python.exe;macOS/Linux 使用 .venv/bin/python
.venv\Scripts\python.exe -m pip install "datasets==4.8.3"

下文运行命令按 Windows 写法给出。使用 macOS 或 Linux 时,把解释器路径换为 .venv/bin/python;后续版本请先用同一份玩具数据复核,再替换生产数据。

本例的竖线 | 是人为约定的分隔符,不是自然语言句子识别器。两条虚构文档各拆成 2 个和 3 个片段,文档标签只是演示可复制字段;真实任务中标签是否适用于每个片段必须另行判断。

同步生成所有输出列

将以下代码存为 map_demo.py,运行 .venv\Scripts\python.exe map_demo.py。Datasets 处理文档说明 batched=True、batch_size 和 remove_columns;这里一次读取两条记录,删除旧列后保留函数重新生成的列。

from datasets import Dataset

ds = Dataset.from_list([
    {'id': 'doc-a', 'text': '安装完成|读取样本', 'label': 1},
    {'id': 'doc-b', 'text': '字段缺失|类型错误|重新核对', 'label': 0},
])

def expand(batch):
    out = {'source_id': [], 'chunk_index': [], 'text': [], 'label': []}
    for doc_id, text, label in zip(batch['id'], batch['text'], batch['label']):
        chunks = [chunk.strip() for chunk in text.split('|') if chunk.strip()]
        for index, chunk in enumerate(chunks):
            out['source_id'].append(doc_id)
            out['chunk_index'].append(index)
            out['text'].append(chunk)
            out['label'].append(label)
    assert len({len(values) for values in out.values()}) == 1
    return out

expanded = ds.map(
    expand, batched=True, batch_size=2, remove_columns=ds.column_names
)
assert len(expanded) == 5
assert list(expanded['source_id']) == ['doc-a', 'doc-a', 'doc-b', 'doc-b', 'doc-b']
assert list(expanded['label']) == [1, 1, 0, 0, 0]
assert list(expanded['chunk_index']) == [0, 1, 0, 1, 2]
for row in expanded:
    print(row)
print('5 rows, aligned columns and source IDs: passed')

remove_columns=ds.column_names 让原始 id、text、label 不再强制保留旧的两行。函数输出重新构造 text 和 label,并增加 source_id 与 chunk_index。列名一样的输出仍会被保留;删除参数不能替你补齐遗漏字段。

batch_size 指输入批次的大小,和训练时的 batch 无关;它也不限制输出行数。最后一个输入批次可能不足两条,函数应同样能处理。这里每次构造一个新的 out,避免不同批次共用残留列表。

核对 5 行输出及来源

实际执行得到 5 条记录,source_id 顺序为 doc-a、doc-a、doc-b、doc-b、doc-b;label 顺序为 1、1、0、0、0;chunk_index 为 0、1、0、1、2。代码逐项断言,确认每个片段的标签和来源没有错位。

先打印少量展开记录,回查原始文档及片段顺序。真实流程应保证 doc_id 全局唯一;如果按多个文件生成局部编号,可把源文件和局部 ID 组合成稳定标识。

仍报列长度错误时怎么排查

  1. 打印 out 中每列的 len,确定返回列本身是否等长;常见问题是追加了文本,却忘记追加标签或来源 ID。
  2. 检查 ds.column_names 里是否存在没有删除的旧列,例如 title、timestamp 或 metadata;新增片段后,它们也必须按片段同步展开或删除。
  3. 检查空片段过滤是否同时影响所有列,不要先独立过滤 text 后再复制全部标签。
  4. 先用 batch_size=2 和一个小样本定位错误,再增加批大小;调大批次不能修复字段对应错误。

函数遇到完全由空片段组成的文本会输出零个片段;输入输出行数不要求相等,要记录每篇文档的保留片段数,防止某些文档意外消失。样例里每列输出长度的断言只证明结构一致,不能证明片段保留了完整语义。

问答:可以把文档标签直接复制到每个片段吗?

只有任务定义允许时才可以。文档整体是正面评价,不代表其中每个句子都正面;如果片段要用于句级分类,可能需要重新标注。拆分后按 source_id 做训练与验证的来源隔离,避免同一文档的不同片段落入两边。

本例验证了行数扩展和列对应,没有测量吞吐速度,也没有证明短片段比原文更适合某个模型。

Ai菜鸟网。发布者:AI小管家,转载请注明出处:https://www.alyyhw.com/30300.html

赞 (0)
AI小管家的头像AI小管家
AI 训练图片怎么找近似重复?用感知哈希生成复核清单
上一篇 1天前
决策树模型怎么生成?用 scikit-learn 训练并导出规则
下一篇 1天前

相关推荐

联系我们

联系我们

1

在线咨询: QQ交谈

邮件:admin@example.com

工作时间:周一至周五,9:30-18:30,节假日休息

关注微信
关注微信
分享本页
返回顶部