File size: 8,043 Bytes
3e936b2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
from typing import Tuple
from einops import rearrange
from torch import nn
import torch.distributed as dist
import torch
from pipeline import SelfForcingTrainingPipeline
from utils.loss import get_denoising_loss
from utils.wan_wrapper import WanDiffusionWrapper, WanTextEncoder, WanVAEWrapper

class BaseModel(nn.Module):

    def __init__(self, args, device):
        super().__init__()
        self._initialize_models(args, device)
        self.device = device
        self.args = args
        self.dtype = torch.bfloat16 if args.mixed_precision else torch.float32
        if hasattr(args, 'denoising_step_list'):
            self.denoising_step_list = torch.tensor(args.denoising_step_list, dtype=torch.long)
            if args.warp_denoising_step:
                timesteps = torch.cat((self.scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32)))
                self.denoising_step_list = timesteps[1000 - self.denoising_step_list]

    def _initialize_models(self, args, device):
        self.real_model_name = getattr(args, 'real_name', 'Wan2.1-T2V-1.3B')
        self.fake_model_name = getattr(args, 'fake_name', 'Wan2.1-T2V-1.3B')
        self.local_attn_size = getattr(args, 'model_kwargs', {}).get('local_attn_size', -1)
        model_kwargs = dict(getattr(args, 'model_kwargs', {}))
        self.generator = WanDiffusionWrapper(**model_kwargs, is_causal=True)
        self.generator.model.requires_grad_(True)
        self.real_score = WanDiffusionWrapper(model_name=self.real_model_name, is_causal=False)
        self.real_score.model.requires_grad_(False)
        self.fake_score = WanDiffusionWrapper(model_name=self.fake_model_name, is_causal=False)
        self.fake_score.model.requires_grad_(True)
        self.text_encoder = WanTextEncoder()
        self.text_encoder.requires_grad_(False)
        self.vae = WanVAEWrapper()
        self.vae.requires_grad_(False)
        self.scheduler = self.generator.get_scheduler()
        self.scheduler.timesteps = self.scheduler.timesteps.to(device)

    def _get_timestep(self, min_timestep: int, max_timestep: int, batch_size: int, num_frame: int, num_frame_per_block: int, uniform_timestep: bool=False) -> torch.Tensor:
        if uniform_timestep:
            timestep = torch.randint(min_timestep, max_timestep, [batch_size, 1], device=self.device, dtype=torch.long).repeat(1, num_frame)
            return timestep
        else:
            timestep = torch.randint(min_timestep, max_timestep, [batch_size, num_frame], device=self.device, dtype=torch.long)
            if self.independent_first_frame:
                timestep_from_second = timestep[:, 1:]
                timestep_from_second = timestep_from_second.reshape(timestep_from_second.shape[0], -1, num_frame_per_block)
                timestep_from_second[:, :, 1:] = timestep_from_second[:, :, 0:1]
                timestep_from_second = timestep_from_second.reshape(timestep_from_second.shape[0], -1)
                timestep = torch.cat([timestep[:, 0:1], timestep_from_second], dim=1)
            else:
                timestep = timestep.reshape(timestep.shape[0], -1, num_frame_per_block)
                timestep[:, :, 1:] = timestep[:, :, 0:1]
                timestep = timestep.reshape(timestep.shape[0], -1)
            return timestep

class SelfForcingModel(BaseModel):

    def __init__(self, args, device):
        super().__init__(args, device)
        self.denoising_loss_func = get_denoising_loss(args.denoising_loss_type)()

    def _run_generator(self, image_or_video_shape, conditional_dict: dict, initial_latent: torch.tensor=None, slice_last_frames: int=21) -> Tuple[torch.Tensor, torch.Tensor]:
        assert getattr(self.args, 'backward_simulation', True), 'Backward simulation needs to be enabled'
        if initial_latent is not None:
            conditional_dict['initial_latent'] = initial_latent
        if self.args.i2v:
            noise_shape = [image_or_video_shape[0], image_or_video_shape[1] - 1, *image_or_video_shape[2:]]
        else:
            noise_shape = image_or_video_shape.copy()
        min_num_frames = self.min_num_training_frames - 1 if self.args.independent_first_frame else self.min_num_training_frames
        max_num_frames = self.num_training_frames - 1 if self.args.independent_first_frame else self.num_training_frames
        assert max_num_frames % self.num_frame_per_block == 0
        assert min_num_frames % self.num_frame_per_block == 0
        max_num_blocks = max_num_frames // self.num_frame_per_block
        min_num_blocks = min_num_frames // self.num_frame_per_block
        num_generated_blocks = torch.randint(min_num_blocks, max_num_blocks + 1, (1,), device=self.device)
        dist.broadcast(num_generated_blocks, src=0)
        num_generated_blocks = num_generated_blocks.item()
        num_generated_frames = num_generated_blocks * self.num_frame_per_block
        if self.args.independent_first_frame and initial_latent is None:
            num_generated_frames += 1
            min_num_frames += 1
        noise_shape[1] = num_generated_frames
        pred_image_or_video, denoised_timestep_from, denoised_timestep_to = self._consistency_backward_simulation(noise=torch.randn(noise_shape, device=self.device, dtype=self.dtype), slice_last_frames=slice_last_frames, **conditional_dict)
        if slice_last_frames != -1 and pred_image_or_video.shape[1] > slice_last_frames:
            with torch.no_grad():
                if slice_last_frames > 1:
                    latent_to_decode = pred_image_or_video[:, :-(slice_last_frames - 1), ...]
                else:
                    latent_to_decode = pred_image_or_video
                pixels = self.vae.decode_to_pixel(latent_to_decode)
                frame = pixels[:, -1:, ...].to(self.dtype)
                frame = rearrange(frame, 'b t c h w -> b c t h w')
                image_latent = self.vae.encode_to_latent(frame).to(self.dtype)
            if slice_last_frames > 1:
                last_frames = pred_image_or_video[:, -(slice_last_frames - 1):, ...]
                pred_image_or_video_sliced = torch.cat([image_latent, last_frames], dim=1)
            else:
                pred_image_or_video_sliced = image_latent
        else:
            pred_image_or_video_sliced = pred_image_or_video
        if num_generated_frames != min_num_frames:
            gradient_mask = torch.ones_like(pred_image_or_video_sliced, dtype=torch.bool)
            if self.args.independent_first_frame:
                gradient_mask[:, :1] = False
            else:
                gradient_mask[:, :self.num_frame_per_block] = False
        else:
            gradient_mask = None
        pred_image_or_video_sliced = pred_image_or_video_sliced.to(self.dtype)
        return (pred_image_or_video_sliced, gradient_mask, denoised_timestep_from, denoised_timestep_to)

    def _consistency_backward_simulation(self, noise: torch.Tensor, slice_last_frames: int=21, **conditional_dict: dict) -> torch.Tensor:
        if self.inference_pipeline is None:
            self._initialize_inference_pipeline()
        return self.inference_pipeline.inference_with_trajectory(noise=noise, **conditional_dict, slice_last_frames=slice_last_frames)

    def _initialize_inference_pipeline(self):
        local_attn_size = getattr(self.args, 'model_kwargs', {}).get('local_attn_size', -1)
        slice_last_frames = getattr(self.args, 'slice_last_frames', 21)
        num_training_frames = getattr(self.args, 'num_training_frames')
        self.inference_pipeline = SelfForcingTrainingPipeline(denoising_step_list=self.denoising_step_list, scheduler=self.scheduler, generator=self.generator, num_frame_per_block=self.num_frame_per_block, independent_first_frame=self.args.independent_first_frame, same_step_across_blocks=self.args.same_step_across_blocks, last_step_only=self.args.last_step_only, num_max_frames=num_training_frames, context_noise=self.args.context_noise, local_attn_size=local_attn_size, slice_last_frames=slice_last_frames, num_training_frames=num_training_frames)