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

Hàm Mất Mát Căn Chỉnh Thời Gian CTC (Connectionist Temporal Classification)

1. Thêm một ký tự trống đặc biệt blank (thường index = 0).

AINâng cao45 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ủ đề

speech-recognitionctc-losssequence-alignmentocrpytorch

Kiến thức tiên quyết: softmax-cross-entropy-from-scratch.

Nội dung đề bài

Mục tiêu kiến thức

  • Trong bài toán Nhận dạng giọng nói (ASR) và Nhận dạng chữ viết (OCR), chuỗi âm thanh đầu vào có độ dài T (ví dụ 1000 frames) lớn hơn rất nhiều so với chuỗi ký tự nhãn U (ví dụ 20 chữ cái), và ta không có căn chỉnh thời gian chi tiết (No ground-truth time alignment).
  • Thuật toán CTC (Graves et al., ICML 2006):
  • Thêm một ký tự trống đặc biệt blank (thường index = 0).
  • Mở rộng chuỗi nhãn mục tiêu bằng cách chèn blank xen kẽ: target = [c1, c2] → l' = [ε, c1, ε, c2, ε] có độ dài 2U + 1.
  • Dùng Quy hoạch động (Forward Variable αt(s)) để tính tổng xác suất của tất cả các đường căn chỉnh hợp lệ dẫn đến nhãn mục tiêu.
  • Hàm mất mát CTC là âm log xác suất tổng này:

LCTC = -ln P(target | x)

Yêu cầu

Viết hàm:

def compute_ctc_loss(
    log_probs: torch.Tensor,
    targets: torch.Tensor,
    input_lengths: torch.Tensor,
    target_lengths: torch.Tensor,
    blank: int = 0
) -> torch.Tensor:
    pass
  • log_probs: Tensor 3D shape (T, B, C) chứa log xác suất sau log_softmax.
  • targets: Tensor 2D shape (B, U) chứa nhãn ký tự.
  • input_lengths: Tensor 1D shape (B,) chứa độ dài chuỗi đầu vào T.
  • target_lengths: Tensor 1D shape (B,) chứa độ dài chuỗi nhãn U.
  • Trả về scalar loss float đối chiếu chuẩn xác với torch.nn.CTCLoss(blank=blank).

Input

  • Hàm compute_ctc_loss(log_probs, targets, input_lengths, target_lengths, blank): 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 compute_ctc_loss: Trả về kết quả kiểu torch.Tensor 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: 3000ms.
  • Giới hạn bộ nhớ: 256MB.
  • 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

T, B, C = (20, 2, 10)
U = 5
logits = torch.randn(T, B, C)
log_probs = F.log_softmax(logits, dim=-1)
targets = torch.randint(1, C, (B, U))
input_lengths = torch.full((B,), T, dtype=torch.long)
target_lengths = torch.full((B,), U, dtype=torch.long)
loss = compute_ctc_loss(log_probs, targets, input_lengths, target_lengths, blank=0)

Output

6.8401

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

T, B, C = (15, 2, 8)
logits = torch.randn(T, B, C)
log_probs = F.log_softmax(logits, dim=-1)
targets = torch.tensor([[1, 2, 3], [4, 5, 0]])
input_lens = torch.tensor([15, 12])
target_lens = torch.tensor([3, 2])
loss = compute_ctc_loss(log_probs, targets, input_lens, target_lens)

Output

9.6971

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.