ai-245Đọc toàn bộ đề miễn phí

Bộ Gom Nhóm Mini-Batch Huấn Luyện Ngôn Ngữ Tự Hồi Quy (Autoregressive Batch Collator)

Khi xây dựng DataLoader để nạp dữ liệu cho mô hình ngôn ngữ tự hồi quy, hàm gom nhóm collate_fn có nhiệm vụ chuyển đổi danh sách các mẫu văn bản độ dài tự do thành một batch hoàn c…

AITrung bình30 phút

Tiến độ của tôi ở bài này

Điểm được lưu vào tài khoản sau khi chấm bài.

Đang tải điểm của bạn…

Kiến thức và chủ đề

llmcollatordataloaderbatching

Kiến thức tiên quyết: causal-next-token-target-shifting, sequence-padding-and-attention-mask.

Nội dung đề bài

Mô tả bài toán

Khi xây dựng DataLoader để nạp dữ liệu cho mô hình ngôn ngữ tự hồi quy, hàm gom nhóm collate_fn có nhiệm vụ chuyển đổi danh sách các mẫu văn bản độ dài tự do thành một batch hoàn chỉnh sẵn sàng cho GPU.

Quy trình kỹ thuật cho mỗi batch:

  • Mỗi mẫu si là một chuỗi token có độ dài Li ≥ 2. Cắt thành:
  • xi = si[:-1] (đầu vào)
  • yi = si[1:] (mục tiêu)
  • Xác định độ dài batch lớn nhất Tmax = maxi len(xi).
  • Đệm xi về Tmax bằng giá trị pad_id.
  • Đệm yi về Tmax bằng giá trị -100 (giá trị mặc định ignore_index trong PyTorch CrossEntropyLoss để bỏ qua việc tính gradient trên các token đệm).
  • Tạo attention_mask shape (B, Tmax): giá trị 1 tại các vị trí token thật của xi, giá trị 0 tại các vị trí đệm pad_id.

Yêu cầu kỹ thuật:

Viết hàm:

def collate_causal_lm_batch(
    samples: list[list[int] | np.ndarray],
    pad_id: int = 0
) -> dict[str, np.ndarray]
  • Nếu samples rỗng: trả về dict rỗng các mảng shape (0, 0).
  • Nếu có bất kỳ mẫu nào có độ dài < 2: ném ngoại lệ ValueError("All samples must have length at least 2").
  • Trả về từ điển:
  • "input_ids": mảng shape (B, Tmax), dtype np.int64.
  • "labels": mảng shape (B, Tmax), dtype np.int64 (đệm bằng -100).
  • "attention_mask": mảng shape (B, Tmax), dtype np.int64.

Input

  • Hàm collate_causal_lm_batch(samples, pad_id): Các tham số đầu vào chứa dữ liệu Tensor/mảng NumPy hoặc giá trị siêu tham số tương ứng.

Output

  • Hàm collate_causal_lm_batch: Trả về kết quả kiểu dict[str, np.ndarray] theo đúng đặc tả kỹ thuật và kích thước quy định.

Ràng buộc

  • Thời gian chạy tối đa: 6000ms.
  • Giới hạn bộ nhớ: 512MB.
  • Dữ liệu đầu vào hợp lệ theo đúng kiểu dữ liệu và miền giá trị được mô tả.

Ví dụ 1

Input

seq1 = [10, 20, 30, 40]
seq2 = [1, 2, 3]
batch = collate_causal_lm_batch([seq1, seq2], pad_id=0)

Output

{
  'input_ids': [[10, 20, 30],
 [ 1,  2,  0]],
  'labels': [[  20,   30,   40],
 [   2,    3, -100]],
  'attention_mask': [[1, 1, 1],
 [1, 1, 0]]
}

Giải thích

Hàm/lớp được gọi với các tham số mẫu trên và trả về kết quả số học / kích thước tensor tương ứng theo đúng thiết kế.

Ví dụ 2

Input

batch = collate_causal_lm_batch([])

Output

{
  'input_ids': [],
  'labels': [],
  'attention_mask': []
}

Giải thích

Hàm/lớp được gọi với các tham số mẫu trên và trả về kết quả số học / kích thước tensor tương ứng theo đúng thiết kế.

3 cấp độ gợi ýMở dần khi bạn thật sự cần hỗ trợ.
Phân tích lời giảiGiải thích hướng tư duy và thuật toán.
Code tham khảoDùng để đối chiếu sau khi tự làm.

Gợi ý và lời giải chỉ mở sau khi bạn bấm Nộp bài. Giáo viên và quản trị viên mở được ngay.