File size: 2,839 Bytes
cfaf85b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Equal-grid Multi-Trace objective and 1..K sampling from Untitled0.ipynb."""
import torch
from trace_imf import TraceIMF


class MultiTraceIMF(TraceIMF):
    def __init__(self, max_nfe=5, jvp_precision="bf16", channels_last=True, **kwargs):
        super().__init__(**kwargs)
        if max_nfe < 1 or jvp_precision not in ("fp32", "bf16"):
            raise ValueError("Invalid max_nfe or JVP precision")
        self.max_nfe = max_nfe
        self.jvp_precision = jvp_precision
        self.channels_last = channels_last

    def draw_interval(self, device, generator=None, nfe=None):
        # One K and j for the entire effective batch, including all microbatches.
        k = int(torch.randint(1, self.max_nfe + 1, (), device=device, generator=generator)) if nfe is None else nfe
        if not 1 <= k <= self.max_nfe:
            raise ValueError("NFE outside the trained range")
        j = int(torch.randint(k, (), device=device, generator=generator))
        return k, j

    def loss(self, model, x, labels=None, *, interval=None, nfe=None, generator=None,
             t=None, r=None, noise=None, derivative_model=None):
        if t is None:
            k, j = interval or self.draw_interval(x.device, generator, nfe)
            r = torch.full((len(x),), j / k, device=x.device)
            t = r + torch.rand(len(x), device=x.device, generator=generator) / k
        elif r is None:
            raise ValueError("Explicit Multi-Trace t requires explicit r")
        return super().loss(model, x, labels, derivative_model=derivative_model,
                            t=t, r=r, noise=noise, generator=generator)

    @torch.no_grad()
    def sample(self, model, n_samples=None, labels=None, *, device=None, noise=None, nfe=1):
        if not 1 <= nfe <= self.max_nfe:
            raise ValueError("NFE outside the trained range")
        device = device or next(model.parameters()).device
        if labels is None:
            raise ValueError("Class-conditioned sampling requires labels")
        labels = labels.to(device)
        z = (torch.randn(len(labels), self.channels, self.image_size, self.image_size, device=device)
             if noise is None else noise.to(device).clone())
        if len(z) != len(labels):
            raise ValueError("Noise and labels must have the same batch size")
        if self.channels_last:
            z = z.contiguous(memory_format=torch.channels_last)
        was_training = model.training
        model.eval()
        try:
            for index in range(nfe, 0, -1):
                t = torch.full((len(z),), index / nfe, device=device)
                r = torch.full_like(t, (index - 1) / nfe)
                z = z - model(z, t, r, y=labels, use_flash_attention=False) / nfe
            return (z.clamp(-1, 1) + 1) * 0.5
        finally:
            model.train(was_training)