QUESTION

zip()怎样组合数据,zip(*data)中的*又做了什么?

A

先理解zip按位置配对,再理解星号把外层序列拆成多个位置参数。

DETAILED ANSWER

直接结论

zip()把多个可迭代对象中相同位置的元素组合在一起;调用函数时的*会把一个序列拆成多个位置参数。因此,zip(*data)就是先把data的每一项分别传给zip,再按位置重新组合,常用来把“多行记录”拆成“多列数据”。

zip怎样按位置组合

names = ["小熊", "小猫"]
scores = [90, 80]

result = zip(names, scores)
print(list(result))

结果:

[("小熊", 90), ("小猫", 80)]

配对过程是:

names[0]与scores[0] → ("小熊", 90)
names[1]与scores[1] → ("小猫", 80)

三个序列也一样:

for doc, title_vector, content_vector in zip(
    DOCS,
    title_vectors,
    content_vectors,
):
    ...

每轮会同时取出同一位置的一篇文档、标题向量和正文向量。

星号不是乘法,而是位置参数解包

假设函数调用是:

values = [1, 2, 3]
func(*values)

它等价于:

func(1, 2, 3)

所以:

zip(*data)

如果data有两项,就等价于:

zip(data[0], data[1])

zip(*data)为什么像行列转换

假设DataLoader收集到一个批次:

data = [
    ("第一段文本", 0),
    ("第二段文本", 1),
    ("第三段文本", 0),
]

执行:

texts, labels = zip(*data)

第一步,*data把三行记录拆成三个参数:

zip(
    ("第一段文本", 0),
    ("第二段文本", 1),
    ("第三段文本", 0),
)

第二步,zip按位置组合:

texts  = ("第一段文本", "第二段文本", "第三段文本")
labels = (0, 1, 0)

原来每行是(文本, 标签),转换后得到“一整列文本”和“一整列标签”。

两个容易忽略的细节

第一,zip()返回的是可迭代对象,不是直接生成好的列表。需要查看全部结果时可以使用list(...)

第二,如果输入长度不同,普通zip()会在最短的输入耗尽时停止:

list(zip([1, 2, 3], ["a", "b"]))
# [(1, "a"), (2, "b")]

因此,文档数量和向量数量本应一一对应时,最好先确认它们的长度相同,避免最后几项被静默忽略。