File size: 2,525 Bytes
f0395ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# Copyright (c) 2025 Hansheng Chen

import torch

from copy import deepcopy
from accelerate import init_empty_weights
from mmgen.models.builder import MODELS, build_module

from .base import BaseModel
from lakonlab.utils import clone_params, tie_untrained_submodules


@MODELS.register_module()
class Diffusion2D(BaseModel):

    def __init__(self,
                 diffusion=dict(type='GMFlow'),
                 diffusion_use_ema=False,
                 tie_ema=True,
                 inference_only=False,
                 train_cfg=None,
                 test_cfg=None):
        super().__init__()
        diffusion.update(train_cfg=train_cfg, test_cfg=test_cfg)
        self.diffusion = build_module(diffusion)
        self.diffusion_use_ema = diffusion_use_ema
        if self.diffusion_use_ema:
            if inference_only:
                self.diffusion_ema = self.diffusion
            else:
                diffusion_ema = deepcopy(diffusion)
                if isinstance(diffusion_ema.get('denoising', None), dict):
                    diffusion_ema['denoising'].pop('pretrained', None)
                with init_empty_weights():
                    self.diffusion_ema = build_module(diffusion_ema)
                if tie_ema:
                    tie_untrained_submodules(self.diffusion_ema, self.diffusion)
                clone_params(self.diffusion_ema, self.diffusion)

        self.train_cfg = dict() if train_cfg is None else deepcopy(train_cfg)
        self.test_cfg = dict() if test_cfg is None else deepcopy(test_cfg)

    def train_minibatch(self, data, loss_scaler=None, running_status=None):
        bs = data['x'].size(0)

        loss, log_vars = self.diffusion(
            data['x'].reshape(bs, 2, 1, 1),
            return_loss=True)

        loss.backward() if loss_scaler is None else loss_scaler.scale(loss).backward()

        return log_vars, bs

    def val_step(self, data, test_cfg_override=dict(), **kwargs):
        bs = data['x'].size(0)
        cfg = deepcopy(self.test_cfg)
        cfg.update(test_cfg_override)
        diffusion = self.diffusion_ema if self.diffusion_use_ema else self.diffusion

        with torch.no_grad():
            if 'noise' in data:
                noise = data['noise'].reshape(bs, 2, 1, 1)
            else:
                noise = torch.randn((bs, 2, 1, 1), device=data['x'].device)
            x_out = diffusion(
                noise=noise,
                test_cfg_override=test_cfg_override)

            return dict(num_samples=bs, pred_x=x_out)