Spaces:
Running
Running
File size: 7,745 Bytes
17f1f54 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 | # coding: utf-8
import torch
import torch.nn.functional as F
from typing import Any, Tuple, List
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
from constants import TARGET_PAD
class Batch:
"""
Batch wrapper that normalizes different batch formats.
Accepts either:
- legacy torchtext batch object with attributes:
.src (tensor), .src_lengths (tensor or int), .trg (tensor), .file_paths (list)
- modern tuple from DataLoader/collate_fn:
(src_padded, src_lengths, trg_padded, files)
Args:
torch_batch: batch object or tuple from DataLoader
pad_index: integer index used for source padding
model: model instance (used to check use_cuda)
"""
def __init__(self, torch_batch: Any, pad_index: int, model: Any):
# Initialize common attributes
self.src = None # LongTensor [B, S]
self.src_lengths = None # LongTensor [B]
self.src_mask = None # Bool/ByteTensor [B,1,S]
self.nseqs = 0
self.trg = None # FloatTensor [B, T, trg_size]
self.trg_input = None # FloatTensor [B, T, trg_size] (for teacher forcing / model input)
self.trg_mask = None # Bool/ByteTensor [B,1,T] (True where not padded)
self.trg_lengths = None # int or LongTensor
self.ntokens = 0 # number of non-pad target frames (sum)
self.file_paths: List[str] = []
# model flags
self.use_cuda = getattr(model, "use_cuda", False)
self.target_pad = TARGET_PAD
# Unpack depending on batch format
# Handle legacy torchtext-like object with attributes
if hasattr(torch_batch, "src") and hasattr(torch_batch, "file_paths"):
# torchtext-style Example batch
try:
# src may be (src_tensor, src_lengths) or src_tensor alone
if isinstance(torch_batch.src, tuple) or isinstance(torch_batch.src, list):
self.src, self.src_lengths = torch_batch.src
else:
self.src = torch_batch.src
# try to obtain lengths if available on batch
self.src_lengths = getattr(torch_batch, "src_lengths", torch.tensor([s.size(0) for s in self.src], dtype=torch.long))
except Exception:
# fallback: assume src is tensor and compute lengths from padding
self.src = torch_batch.src
self.src_lengths = getattr(torch_batch, "src_lengths", torch.sum(self.src != pad_index, dim=1))
self.file_paths = list(getattr(torch_batch, "file_paths", []))
# Targets (if present)
if hasattr(torch_batch, "trg"):
self.trg = torch_batch.trg
# Handle tuple produced by DataLoader / collate_fn
elif isinstance(torch_batch, (tuple, list)) and len(torch_batch) >= 3:
# Expected format: (src_padded, src_lengths, trg_padded, files)
# Some collate_fns may not return src_lengths; handle both cases.
# Common expected:
# src_padded: LongTensor [B, S]
# src_lengths: LongTensor [B]
# trg_padded: FloatTensor [B, T, trg_size]
# files: list[str]
try:
self.src = torch_batch[0]
self.src_lengths = torch_batch[1]
self.trg = torch_batch[2]
# files may be absent or None
if len(torch_batch) > 3:
self.file_paths = list(torch_batch[3])
else:
self.file_paths = []
except Exception:
raise ValueError("Unrecognized tuple batch format. Expected (src, src_lengths, trg, files).")
else:
raise ValueError("Unrecognized batch format passed to Batch.")
# Ensure shapes / dtypes
if isinstance(self.src_lengths, int):
self.src_lengths = torch.tensor([self.src_lengths] * self.src.size(0), dtype=torch.long)
if self.src is not None and not isinstance(self.src_lengths, torch.Tensor):
# attempt to compute lengths from padding if possible
try:
self.src_lengths = torch.sum(self.src != pad_index, dim=1).to(torch.long)
except Exception:
# fallback zeros
self.src_lengths = torch.zeros(self.src.size(0), dtype=torch.long)
# src_mask: True where token != pad_index
if self.src is not None:
self.src_mask = (self.src != pad_index).unsqueeze(1) # [B,1,S]
self.nseqs = self.src.size(0)
# Targets handling
if self.trg is not None:
# trg is expected shape [B, T, trg_size]
# trg_lengths: number of frames (T) - if not available infer from shape
try:
self.trg_lengths = self.trg.shape[1]
except Exception:
self.trg_lengths = None
# trg_input: the model expects target input. Keep same shape as trg.
# If you want to shift / remove last frame, do that in the training loop or here by uncommenting the next line:
# self.trg_input = self.trg[:, :-1, :].clone()
self.trg_input = self.trg.clone()
# trg mask: True where frame is not padding. We assume padding frames have all elements equal to TARGET_PAD.
# To detect padded frames, compare first element of each frame to TARGET_PAD (fast and typical).
# Fallback: compare sum across frame (if consistent).
try:
# If last dim exists
if self.trg.dim() == 3:
self.trg_mask = (self.trg[:, :, 0] != self.target_pad).unsqueeze(1) # [B,1,T]
else:
# unexpected dimensions -> assume all valid
self.trg_mask = torch.ones((self.trg.size(0), 1, self.trg.size(1)), dtype=torch.bool)
except Exception:
# fallback: assume all frames valid
self.trg_mask = torch.ones((self.trg.size(0), 1, self.trg.size(1)), dtype=torch.bool)
# ntokens: number of non-padding frames across the batch
try:
self.ntokens = int(torch.sum(self.trg_mask).item())
except Exception:
self.ntokens = 0
# Filepaths: ensure list
if self.file_paths is None:
self.file_paths = []
# Move to GPU if required
if self.use_cuda:
self._make_cuda()
def _make_cuda(self):
"""Move the batch tensors to GPU (in-place)."""
if self.src is not None:
self.src = self.src.to(device)
self.src_mask = self.src_mask.to(device)
self.src_lengths = self.src_lengths.to(device)
if self.trg_input is not None:
self.trg_input = self.trg_input.to(device)
if self.trg is not None:
self.trg = self.trg.to(device)
if self.trg_mask is not None:
self.trg_mask = self.trg_mask.to(device)
def to(self, device_: torch.device):
"""Move batch to specified device and return self (convenience)."""
if self.src is not None:
self.src = self.src.to(device_)
self.src_mask = self.src_mask.to(device_)
self.src_lengths = self.src_lengths.to(device_)
if self.trg_input is not None:
self.trg_input = self.trg_input.to(device_)
if self.trg is not None:
self.trg = self.trg.to(device_)
if self.trg_mask is not None:
self.trg_mask = self.trg_mask.to(device_)
return self
|