0Pricing
Deep Learning Academy · 课时

使用 collate_fn 处理可变长度输入

填充并堆叠长度不齐的样本

使用 collate_fn 处理可变长度输入 是 CoddyKit 上的免费 Deep Learning Academy 课时。 这是第 3 节课,共 4 节。 你可以在下方免费阅读本课时的完整内容 — 然后在浏览器中使用内置代码编辑器和全天候 AI 导师进行实践。 这是 Deep Learning Academy 学习路径的一部分,你的进度在网页和 CoddyKit 应用中同步。 Deep Learning Academy 课程共包含 4 节课。

本课时的部分内容尚未翻译,以英文显示。

When Samples Don't Match

Stacking into a batch needs every sample the same shape. But sentences and audio clips have different lengths, so the default collate step fails. 🧩

What collate_fn Does

The DataLoader gathers a list of samples and passes them to collate_fn, which merges them into one batch. By default it simply stacks tensors.

Ragged Inputs Break Stacking

Try to stack a length-5 and a length-8 sequence and PyTorch raises a shape error. Ragged lengths are exactly the case a custom collate must handle.

Write Your Own collate_fn

You pass a function to the DataLoader's collate_fn argument. It receives a list of samples and returns whatever batch shape your model expects.

loader = DataLoader(ds, batch_size=4, collate_fn=my_collate)

Step One: Split the List

Inside your function, unzip the list of pairs into separate sequences and labels. Now you can treat each group on its own before merging.

def my_collate(batch):
    seqs, labels = zip(*batch)

Pad to the Longest

The trick for variable lengths is padding: extend every sequence to the longest one with a filler value, so they finally share a shape.

pad_sequence Does It for You

PyTorch ships pad_sequence, which pads a list of tensors to equal length and stacks them. Set batch_first so the batch dimension comes first.

from torch.nn.utils.rnn import pad_sequence
padded = pad_sequence(seqs, batch_first=True)

Remember the Real Lengths

Padding adds fake tokens, so also return each sequence's true length. Your model uses these to ignore the padded positions during the forward pass.

lengths = torch.tensor([len(s) for s in seqs])

Stack the Labels

Labels are usually fixed size, so a normal stack works for them. Return the padded inputs, the lengths, and the stacked labels together.

labels = torch.stack(labels)
return padded, lengths, labels

Mask Out the Padding

Later you build a mask from the lengths so the loss and attention skip padded slots. Padding fills shape without polluting the gradients.

One Function, Any Shape

With a custom collate_fn, the same DataLoader handles text, audio, and graphs. You control exactly how loose samples become one tidy batch.

Quick Check

Why do variable-length sequences need a custom collate_fn?

Recap

A custom collate_fn turns a list of uneven samples into one batch, usually by padding sequences to equal length and tracking their real sizes. 🎉

常见问题解答

「使用 collate_fn 处理可变长度输入」课时是免费的吗?

是的 — 「使用 collate_fn 处理可变长度输入」的完整文本可在网页上免费阅读。要进行交互式练习(内置代码编辑器和全天候 AI 导师)并解锁 Deep Learning Academy 课程的其余内容,请升级到 CoddyKit PRO。 Deep Learning Academy 课程共包含 4 节课。

「使用 collate_fn 处理可变长度输入」这节课中我会学到什么?

填充并堆叠长度不齐的样本 你通过在浏览器中直接运行的动手代码来练习 Deep Learning Academy,全天候 AI 导师会在你学习这节课的过程中回答你的问题。

学习 Deep Learning Academy 需要有经验吗?

无需任何先前经验。CoddyKit 上的 Deep Learning Academy 课程适合初学者到高级学习者,你可以从这里开始或从头开始,按照自己的节奏学习。 这是第 3 节课,共 4 节。

「使用 collate_fn 处理可变长度输入」课时需要多长时间?

大多数 CoddyKit 课程大约需要 5–10 分钟。每节课都很精短且互动,所以你能稳步进步,并在网页和应用中从离开的地方继续。

我能在这节 Deep Learning Academy 课中编写并运行代码吗?

能。每节 Deep Learning Academy 课都包含内置代码编辑器,你可以在浏览器中直接编写并运行真实代码,并获得即时 AI 反馈 — 无需本地设置。

此课程中的所有课时

  1. 编写自定义数据集类
  2. 批处理、打乱与 num_workers
  3. 使用 collate_fn 处理可变长度输入
  4. 归一化与标准化输入
← 返回 Deep Learning Academy