import torch


# ranges: {batch_idx: [(start, end), ...]}
# pad_idxs: bool tensor (bsz, slen)
def triangular_mask_slow(bsz, slen, ranges, pad_mask, device):
    # returns (bsz, slen, slen)
    # mask[b, i, j] = 0 if tokens[b, i] attends to tokens[b, j], otherwise -inf
    # mask[b, i, i] must be 0
    #
    mask = torch.full((bsz, slen, slen), -float("inf"), device=device)
    for i in range(slen):
        mask[:, i, i] = 0
    for b in range(bsz):
        for rng in ranges[b]:
            for i in range(rng[0], rng[1]):
                for j in range(rng[0], i + 1):
                    if not pad_mask[b, j]:
                        continue
                    mask[b, i, j] = 0
    return mask


def triangular_mask(bsz, slen, ranges, pad_mask, device):
    mask = torch.full((bsz, slen, slen), -float("inf"), device=device)

    for b in range(bsz):
        for rng in ranges[b]:
            mask[b, rng[0] : rng[1], rng[0] : rng[1]].triu_(0)

    mask.masked_fill_(pad_mask.unsqueeze(1) == False, -float("inf"))

    for b in range(bsz):
        mask[b].fill_diagonal_(0.0)

    return mask


def triangular_mask_for_pad(xs, pad_symbol):
    bsz, slen = xs.shape[:2]
    mask = torch.full((bsz, slen, slen), -float("inf"), device=xs.device)
    mask.triu_(0)
    syms = xs[:, :, 0] if xs.ndim == 3 else xs
    mask.masked_fill_((syms == pad_symbol).unsqueeze(1), -float("inf"))
    for b in range(bsz):
        mask[b].fill_diagonal_(0.0)
    return mask
