HarryLee/eCommerceImageCaptioning
2
1# Copyright 2022 The OFA-Sys Team. 2# All rights reserved.3# This source code is licensed under the Apache 2.0 license 4# found in the LICENSE file in the root directory.5 6try:7 from collections.abc import Iterable8except ImportError:9 from collections import Iterable10import contextlib11import itertools12import logging13import re14import warnings15from typing import Optional, Tuple16 17import numpy as np18import torch19 20from fairseq.file_io import PathManager21from fairseq import utils22import os23 24logger = logging.getLogger(__name__)25 26 27def infer_language_pair(path):28 """Infer language pair from filename: <split>.<lang1>-<lang2>.(...).idx"""29 src, dst = None, None30 for filename in PathManager.ls(path):31 parts = filename.split(".")32 if len(parts) >= 3 and len(parts[1].split("-")) == 2:33 return parts[1].split("-")34 return src, dst35 36 37def collate_tokens(38 values,39 pad_idx,40 eos_idx=None,41 left_pad=False,42 move_eos_to_beginning=False,43 pad_to_length=None,44 pad_to_multiple=1,45 pad_to_bsz=None,46):47 """Convert a list of 1d tensors into a padded 2d tensor."""48 size = max(v.size(0) for v in values)49 size = size if pad_to_length is None else max(size, pad_to_length)50 if pad_to_multiple != 1 and size % pad_to_multiple != 0:51 size = int(((size - 0.1) // pad_to_multiple + 1) * pad_to_multiple)52 53 def copy_tensor(src, dst):54 assert dst.numel() == src.numel()55 if move_eos_to_beginning:56 if eos_idx is None:57 # if no eos_idx is specified, then use the last token in src58 dst[0] = src[-1]59 else:60 dst[0] = eos_idx61 dst[1:] = src[:-1]62 else:63 dst.copy_(src)64 65 if values[0].dim() == 1:66 res = values[0].new(len(values), size).fill_(pad_idx)67 elif values[0].dim() == 2:68 assert move_eos_to_beginning is False69 res = values[0].new(len(values), size, values[0].size(1)).fill_(pad_idx)70 else:71 raise NotImplementedError72 73 for i, v in enumerate(values):74 copy_tensor(v, res[i][size - len(v) :] if left_pad else res[i][: len(v)])75 return res76 77 78def load_indexed_dataset(79 path, dictionary=None, dataset_impl=None, combine=False, default="cached"80):81 """A helper function for loading indexed datasets.82 83 Args:84 path (str): path to indexed dataset (e.g., 'data-bin/train')85 dictionary (~fairseq.data.Dictionary): data dictionary86 dataset_impl (str, optional): which dataset implementation to use. If87 not provided, it will be inferred automatically. For legacy indexed88 data we use the 'cached' implementation by default.89 combine (bool, optional): automatically load and combine multiple90 datasets. For example, if *path* is 'data-bin/train', then we will91 combine 'data-bin/train', 'data-bin/train1', ... and return a92 single ConcatDataset instance.93 """94 import fairseq.data.indexed_dataset as indexed_dataset95 from fairseq.data.concat_dataset import ConcatDataset96 97 datasets = []98 for k in itertools.count():99 path_k = path + (str(k) if k > 0 else "")100 try:101 path_k = indexed_dataset.get_indexed_dataset_to_local(path_k)102 except Exception as e:103 if "StorageException: [404] Path not found" in str(e):104 logger.warning(f"path_k: {e} not found")105 else:106 raise e107 108 dataset_impl_k = dataset_impl109 if dataset_impl_k is None:110 dataset_impl_k = indexed_dataset.infer_dataset_impl(path_k)111 dataset = indexed_dataset.make_dataset(112 path_k,113 impl=dataset_impl_k or default,114 fix_lua_indexing=True,115 dictionary=dictionary,116 )117 if dataset is None:118 break119 logger.info("loaded {:,} examples from: {}".format(len(dataset), path_k))120 datasets.append(dataset)121 if not combine:122 break123 if len(datasets) == 0:124 return None125 elif len(datasets) == 1:126 return datasets[0]127 else:128 return ConcatDataset(datasets)129 130 131@contextlib.contextmanager132def numpy_seed(seed, *addl_seeds):133 """Context manager which seeds the NumPy PRNG with the specified seed and134 restores the state afterward"""135 if seed is None:136 yield137 return138 if len(addl_seeds) > 0:139 seed = int(hash((seed, *addl_seeds)) % 1e6)140 state = np.random.get_state()141 np.random.seed(seed)142 try:143 yield144 finally:145 np.random.set_state(state)146 147 148def collect_filtered(function, iterable, filtered):149 """150 Similar to :func:`filter` but collects filtered elements in ``filtered``.151 152 Args:153 function (callable): function that returns ``False`` for elements that154 should be filtered155 iterable (iterable): iterable to filter156 filtered (list): list to store filtered elements157 """158 for el in iterable:159 if function(el):160 yield el161 else:162 filtered.append(el)163 164 165def _filter_by_size_dynamic(indices, size_fn, max_positions, raise_exception=False):166 def compare_leq(a, b):167 return a <= b if not isinstance(a, tuple) else max(a) <= b168 169 def check_size(idx):170 if isinstance(max_positions, float) or isinstance(max_positions, int):171 return size_fn(idx) <= max_positions172 elif isinstance(max_positions, dict):173 idx_size = size_fn(idx)174 assert isinstance(idx_size, dict)175 intersect_keys = set(max_positions.keys()) & set(idx_size.keys())176 return all(177 all(178 a is None or b is None or a <= b179 for a, b in zip(idx_size[key], max_positions[key])180 )181 for key in intersect_keys182 )183 else:184 # For MultiCorpusSampledDataset, will generalize it later185 if not isinstance(size_fn(idx), Iterable):186 return all(size_fn(idx) <= b for b in max_positions)187 return all(188 a is None or b is None or a <= b189 for a, b in zip(size_fn(idx), max_positions)190 )191 192 ignored = []193 itr = collect_filtered(check_size, indices, ignored)194 indices = np.fromiter(itr, dtype=np.int64, count=-1)195 return indices, ignored196 197 198def filter_by_size(indices, dataset, max_positions, raise_exception=False):199 """200 [deprecated] Filter indices based on their size.201 Use `FairseqDataset::filter_indices_by_size` instead.202 203 Args:204 indices (List[int]): ordered list of dataset indices205 dataset (FairseqDataset): fairseq dataset instance206 max_positions (tuple): filter elements larger than this size.207 Comparisons are done component-wise.208 raise_exception (bool, optional): if ``True``, raise an exception if209 any elements are filtered (default: False).210 """211 warnings.warn(212 "data_utils.filter_by_size is deprecated. "213 "Use `FairseqDataset::filter_indices_by_size` instead.",214 stacklevel=2,215 )216 if isinstance(max_positions, float) or isinstance(max_positions, int):217 if hasattr(dataset, "sizes") and isinstance(dataset.sizes, np.ndarray):218 ignored = indices[dataset.sizes[indices] > max_positions].tolist()219 indices = indices[dataset.sizes[indices] <= max_positions]220 elif (221 hasattr(dataset, "sizes")222 and isinstance(dataset.sizes, list)223 and len(dataset.sizes) == 1224 ):225 ignored = indices[dataset.sizes[0][indices] > max_positions].tolist()226 indices = indices[dataset.sizes[0][indices] <= max_positions]227 else:228 indices, ignored = _filter_by_size_dynamic(229 indices, dataset.size, max_positions230 )231 else:232 indices, ignored = _filter_by_size_dynamic(indices, dataset.size, max_positions)233 234 if len(ignored) > 0 and raise_exception:235 raise Exception(236 (237 "Size of sample #{} is invalid (={}) since max_positions={}, "238 "skip this example with --skip-invalid-size-inputs-valid-test"239 ).format(ignored[0], dataset.size(ignored[0]), max_positions)240 )241 if len(ignored) > 0:242 logger.warning(243 (244 "{} samples have invalid sizes and will be skipped, "245 "max_positions={}, first few sample ids={}"246 ).format(len(ignored), max_positions, ignored[:10])247 )248 return indices249 250 251def filter_paired_dataset_indices_by_size(src_sizes, tgt_sizes, indices, max_sizes):252 """Filter a list of sample indices. Remove those that are longer253 than specified in max_sizes.254 255 Args:256 indices (np.array): original array of sample indices257 max_sizes (int or list[int] or tuple[int]): max sample size,258 can be defined separately for src and tgt (then list or tuple)259 260 Returns:261 np.array: filtered sample array262 list: list of removed indices263 """264 if max_sizes is None:265 return indices, []266 if type(max_sizes) in (int, float):267 max_src_size, max_tgt_size = max_sizes, max_sizes268 else:269 max_src_size, max_tgt_size = max_sizes270 if tgt_sizes is None:271 ignored = indices[src_sizes[indices] > max_src_size]272 else:273 ignored = indices[274 (src_sizes[indices] > max_src_size) | (tgt_sizes[indices] > max_tgt_size)275 ]276 if len(ignored) > 0:277 if tgt_sizes is None:278 indices = indices[src_sizes[indices] <= max_src_size]279 else:280 indices = indices[281 (src_sizes[indices] <= max_src_size)282 & (tgt_sizes[indices] <= max_tgt_size)283 ]284 return indices, ignored.tolist()285 286 287def batch_by_size(288 indices,289 num_tokens_fn,290 num_tokens_vec=None,291 max_tokens=None,292 max_sentences=None,293 required_batch_size_multiple=1,294 fixed_shapes=None,295):296 """297 Yield mini-batches of indices bucketed by size. Batches may contain298 sequences of different lengths.299 300 Args:301 indices (List[int]): ordered list of dataset indices302 num_tokens_fn (callable): function that returns the number of tokens at303 a given index304 num_tokens_vec (List[int], optional): precomputed vector of the number305 of tokens for each index in indices (to enable faster batch generation)306 max_tokens (int, optional): max number of tokens in each batch307 (default: None).308 max_sentences (int, optional): max number of sentences in each309 batch (default: None).310 required_batch_size_multiple (int, optional): require batch size to311 be less than N or a multiple of N (default: 1).312 fixed_shapes (List[Tuple[int, int]], optional): if given, batches will313 only be created with the given shapes. *max_sentences* and314 *required_batch_size_multiple* will be ignored (default: None).315 """316 try:317 from fairseq.data.data_utils_fast import (318 batch_by_size_fn,319 batch_by_size_vec,320 batch_fixed_shapes_fast,321 )322 except ImportError:323 raise ImportError(324 "Please build Cython components with: "325 "`python setup.py build_ext --inplace`"326 )327 except ValueError:328 raise ValueError(329 "Please build (or rebuild) Cython components with `python setup.py build_ext --inplace`."330 )331 332 # added int() to avoid TypeError: an integer is required333 max_tokens = (334 int(max_tokens) if max_tokens is not None else -1335 )336 max_sentences = max_sentences if max_sentences is not None else -1337 bsz_mult = required_batch_size_multiple338 339 if not isinstance(indices, np.ndarray):340 indices = np.fromiter(indices, dtype=np.int64, count=-1)341 342 if num_tokens_vec is not None and not isinstance(num_tokens_vec, np.ndarray):343 num_tokens_vec = np.fromiter(num_tokens_vec, dtype=np.int64, count=-1)344 345 if fixed_shapes is None:346 if num_tokens_vec is None:347 return batch_by_size_fn(348 indices,349 num_tokens_fn,350 max_tokens,351 max_sentences,352 bsz_mult,353 )354 else:355 return batch_by_size_vec(356 indices,357 num_tokens_vec,358 max_tokens,359 max_sentences,360 bsz_mult,361 )362 363 else:364 fixed_shapes = np.array(fixed_shapes, dtype=np.int64)365 sort_order = np.lexsort(366 [367 fixed_shapes[:, 1].argsort(), # length368 fixed_shapes[:, 0].argsort(), # bsz369 ]370 )371 fixed_shapes_sorted = fixed_shapes[sort_order]372 return batch_fixed_shapes_fast(indices, num_tokens_fn, fixed_shapes_sorted)373 374 375def post_process(sentence: str, symbol: str):376 if symbol == "sentencepiece":377 sentence = sentence.replace(" ", "").replace("\u2581", " ").strip()378 elif symbol == "wordpiece":379 sentence = sentence.replace(" ", "").replace("_", " ").strip()380 elif symbol == "letter":381 sentence = sentence.replace(" ", "").replace("|", " ").strip()382 elif symbol == "silence":383 import re384 sentence = sentence.replace("<SIL>", "")385 sentence = re.sub(' +', ' ', sentence).strip()386 elif symbol == "_EOW":387 sentence = sentence.replace(" ", "").replace("_EOW", " ").strip()388 elif symbol in {"subword_nmt", "@@ ", "@@"}:389 if symbol == "subword_nmt":390 symbol = "@@ "391 sentence = (sentence + " ").replace(symbol, "").rstrip()392 elif symbol == "none":393 pass394 elif symbol is not None:395 raise NotImplementedError(f"Unknown post_process option: {symbol}")396 return sentence397 398 399def compute_mask_indices(400 shape: Tuple[int, int],401 padding_mask: Optional[torch.Tensor],402 mask_prob: float,403 mask_length: int,404 mask_type: str = "static",405 mask_other: float = 0.0,406 min_masks: int = 0,407 no_overlap: bool = False,408 min_space: int = 0,409) -> np.ndarray:410 """411 Computes random mask spans for a given shape412 413 Args:414 shape: the the shape for which to compute masks.415 should be of size 2 where first element is batch size and 2nd is timesteps416 padding_mask: optional padding mask of the same size as shape, which will prevent masking padded elements417 mask_prob: probability for each token to be chosen as start of the span to be masked. this will be multiplied by418 number of timesteps divided by length of mask span to mask approximately this percentage of all elements.419 however due to overlaps, the actual number will be smaller (unless no_overlap is True)420 mask_type: how to compute mask lengths421 static = fixed size422 uniform = sample from uniform distribution [mask_other, mask_length*2]423 normal = sample from normal distribution with mean mask_length and stdev mask_other. mask is min 1 element424 poisson = sample from possion distribution with lambda = mask length425 min_masks: minimum number of masked spans426 no_overlap: if false, will switch to an alternative recursive algorithm that prevents spans from overlapping427 min_space: only used if no_overlap is True, this is how many elements to keep unmasked between spans428 """429 430 bsz, all_sz = shape431 mask = np.full((bsz, all_sz), False)432 433 all_num_mask = int(434 # add a random number for probabilistic rounding435 mask_prob * all_sz / float(mask_length)436 + np.random.rand()437 )438 439 all_num_mask = max(min_masks, all_num_mask)440 441 mask_idcs = []442 for i in range(bsz):443 if padding_mask is not None:444 sz = all_sz - padding_mask[i].long().sum().item()445 num_mask = int(446 # add a random number for probabilistic rounding447 mask_prob * sz / float(mask_length)448 + np.random.rand()449 )450 num_mask = max(min_masks, num_mask)451 else:452 sz = all_sz453 num_mask = all_num_mask454 455 if mask_type == "static":456 lengths = np.full(num_mask, mask_length)457 elif mask_type == "uniform":458 lengths = np.random.randint(mask_other, mask_length * 2 + 1, size=num_mask)459 elif mask_type == "normal":460 lengths = np.random.normal(mask_length, mask_other, size=num_mask)461 lengths = [max(1, int(round(x))) for x in lengths]462 elif mask_type == "poisson":463 lengths = np.random.poisson(mask_length, size=num_mask)464 lengths = [int(round(x)) for x in lengths]465 else:466 raise Exception("unknown mask selection " + mask_type)467 468 if sum(lengths) == 0:469 lengths[0] = min(mask_length, sz - 1)470 471 if no_overlap:472 mask_idc = []473 474 def arrange(s, e, length, keep_length):475 span_start = np.random.randint(s, e - length)476 mask_idc.extend(span_start + i for i in range(length))477 478 new_parts = []479 if span_start - s - min_space >= keep_length:480 new_parts.append((s, span_start - min_space + 1))481 if e - span_start - keep_length - min_space > keep_length:482 new_parts.append((span_start + length + min_space, e))483 return new_parts484 485 parts = [(0, sz)]486 min_length = min(lengths)487 for length in sorted(lengths, reverse=True):488 lens = np.fromiter(489 (e - s if e - s >= length + min_space else 0 for s, e in parts),490 np.int,491 )492 l_sum = np.sum(lens)493 if l_sum == 0:494 break495 probs = lens / np.sum(lens)496 c = np.random.choice(len(parts), p=probs)497 s, e = parts.pop(c)498 parts.extend(arrange(s, e, length, min_length))499 mask_idc = np.asarray(mask_idc)500 else:501 min_len = min(lengths)502 if sz - min_len <= num_mask:503 min_len = sz - num_mask - 1504 505 mask_idc = np.random.choice(sz - min_len, num_mask, replace=False)506 507 mask_idc = np.asarray(508 [509 mask_idc[j] + offset510 for j in range(len(mask_idc))511 for offset in range(lengths[j])512 ]513 )514 515 mask_idcs.append(np.unique(mask_idc[mask_idc < sz]))516 517 min_len = min([len(m) for m in mask_idcs])518 for i, mask_idc in enumerate(mask_idcs):519 if len(mask_idc) > min_len:520 mask_idc = np.random.choice(mask_idc, min_len, replace=False)521 mask[i, mask_idc] = True522 523 return mask524 525 526def get_mem_usage():527 try:528 import psutil529 530 mb = 1024 * 1024531 return f"used={psutil.virtual_memory().used / mb}Mb; avail={psutil.virtual_memory().available / mb}Mb"532 except ImportError:533 return "N/A"534 535 536# lens: torch.LongTensor537# returns: torch.BoolTensor538def lengths_to_padding_mask(lens):539 bsz, max_lens = lens.size(0), torch.max(lens).item()540 mask = torch.arange(max_lens).to(lens.device).view(1, max_lens)541 mask = mask.expand(bsz, -1) >= lens.view(bsz, 1).expand(-1, max_lens)542 return mask543 544 545# lens: torch.LongTensor546# returns: torch.BoolTensor547def lengths_to_mask(lens):548 return ~lengths_to_padding_mask(lens)549 550 551def get_buckets(sizes, num_buckets):552 buckets = np.unique(553 np.percentile(554 sizes,555 np.linspace(0, 100, num_buckets + 1),556 interpolation='lower',557 )[1:]558 )559 return buckets560 561 562def get_bucketed_sizes(orig_sizes, buckets):563 sizes = np.copy(orig_sizes)564 assert np.min(sizes) >= 0565 start_val = -1566 for end_val in buckets:567 mask = (sizes > start_val) & (sizes <= end_val)568 sizes[mask] = end_val569 start_val = end_val570 return sizes571 572 573 574def _find_extra_valid_paths(dataset_path: str) -> set:575 paths = utils.split_paths(dataset_path)576 all_valid_paths = set()577 for sub_dir in paths:578 contents = PathManager.ls(sub_dir)579 valid_paths = [c for c in contents if re.match("valid*[0-9].*", c) is not None]580 all_valid_paths |= {os.path.basename(p) for p in valid_paths}581 # Remove .bin, .idx etc582 roots = {os.path.splitext(p)[0] for p in all_valid_paths}583 return roots584 585 586def raise_if_valid_subsets_unintentionally_ignored(train_cfg) -> None:587 """Raises if there are paths matching 'valid*[0-9].*' which are not combined or ignored."""588 if (589 train_cfg.dataset.ignore_unused_valid_subsets590 or train_cfg.dataset.combine_valid_subsets591 or train_cfg.dataset.disable_validation592 or not hasattr(train_cfg.task, "data")593 ):594 return595 other_paths = _find_extra_valid_paths(train_cfg.task.data)596 specified_subsets = train_cfg.dataset.valid_subset.split(",")597 ignored_paths = [p for p in other_paths if p not in specified_subsets]598 if ignored_paths:599 advice = "Set --combine-val to combine them or --ignore-unused-valid-subsets to ignore them."600 msg = f"Valid paths {ignored_paths} will be ignored. {advice}"601 raise ValueError(msg)602 