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