import torch from contextlib import contextmanager import time import logging import os from collections import OrderedDict def subsequent_mask(size): "Mask out subsequent positions." attn_shape = (1, size, size) subsequent_mask = torch.triu(torch.ones(attn_shape), diagonal=1).type(torch.uint8) return subsequent_mask == 0 @contextmanager def stopwatch(name): start = time.time() yield logging.info(f"{name}: {1000*(time.time() - start):.1f}ms") def stopwatch_iter(name, iter=10): deltas = [] for _ in range(iter): start = time.time() yield deltas.append(time.time() - start) logging.info(f"{name}: {1000*(sum(deltas) / len(deltas)):.1f}ms") def configure_logging(): logging.getLogger().setLevel( logging.getLevelName(os.getenv("COMPOSER_LOGLEVEL", "INFO").upper()) ) def remove_module_prefix(d): module = "module." return OrderedDict( [(k[len(module) :] if k.startswith(module) else k, v) for k, v in d.items()] ) def add_layer_sequence(d): return OrderedDict( [ (k.replace("transformer.layers", "transformer.layer_sequence.layers"), v) for k, v in d.items() ] )