File size: 2,636 Bytes
3194a67
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
from __future__ import annotations

import torch


class MultiMaskDiffusion:
    """Forward corruption from the MultiMDM paper.

    p_t(x_t|x_0) = alpha_t * delta(x_t=x_0)
      + (1-alpha_t) * [beta_t * delta(x_t=m_x0)
      + (1-beta_t) * Uniform(M masks)].
    """

    def __init__(
        self,
        vocab_size: int,
        clean_token_start: int,
        num_masks: int,
        mask_token_start: int,
        pad_id: int = 0,
    ):
        self.vocab_size = vocab_size
        self.clean_token_start = clean_token_start
        self.num_masks = num_masks
        self.mask_token_start = mask_token_start
        self.pad_id = pad_id

    def designated_mask(self, clean_ids: torch.Tensor) -> torch.Tensor:
        if torch.any(clean_ids < self.clean_token_start):
            raise ValueError("designated_mask expects clean vocabulary token IDs")
        return self.mask_token_start + (
            (clean_ids - self.clean_token_start) % self.num_masks
        )

    @staticmethod
    def alpha(t: torch.Tensor) -> torch.Tensor:
        return (1.0 - t.clamp(0.0, 1.0)).clamp(0.0, 1.0)

    @staticmethod
    def beta(t: torch.Tensor) -> torch.Tensor:
        return (1.0 - t.clamp(0.0, 1.0)).clamp(0.0, 1.0)

    def q_sample(
        self,
        x0: torch.Tensor,
        t: torch.Tensor,
        attention_mask: torch.Tensor | None = None,
    ) -> torch.Tensor:
        if t.ndim == 0:
            t = t.expand(x0.shape[0])
        if t.ndim != 1 or t.shape[0] != x0.shape[0]:
            raise ValueError("t must be scalar or have shape [batch]")
        view_shape = [x0.shape[0]] + [1] * (x0.ndim - 1)
        alpha = self.alpha(t).view(*view_shape)
        beta = self.beta(t).view(*view_shape)

        valid = x0.ge(self.clean_token_start)
        if attention_mask is not None:
            valid = valid & attention_mask.bool()

        keep_clean = torch.rand_like(x0, dtype=torch.float32) < alpha
        use_designated = torch.rand_like(x0, dtype=torch.float32) < beta
        uniform_masks = self.mask_token_start + torch.randint(
            self.num_masks, x0.shape, device=x0.device
        )
        safe_clean = torch.where(valid, x0, torch.full_like(x0, self.clean_token_start))
        designated = self.designated_mask(safe_clean)
        sampled_mask = torch.where(use_designated, designated, uniform_masks)
        corrupted = torch.where(keep_clean, x0, sampled_mask)
        return torch.where(valid, corrupted, x0)

    def is_mask(self, ids: torch.Tensor) -> torch.Tensor:
        return ids.ge(self.mask_token_start) & ids.lt(
            self.mask_token_start + self.num_masks
        )