Team Ai
Apppublic

HarryLee/eCommerceImageCaptioning

sourceHugging Faceupdated 4y agoView on Hugging Face
2likes
data_utils.py602 linesDownload Raw Back to data
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