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()
        ]
    )
