MCPcopy Create free account
hub / github.com/Francis-Rings/StableAnimator / forward

Method forward

animation/modules/refined_vae.py:75–150  ·  view source on GitHub ↗

r"""The forward method of the `Decoder` class.

(
        self,
        sample: torch.Tensor,
        image_only_indicator: torch.Tensor,
        num_frames: int = 1,
    )

Source from the content-addressed store, hash-verified

73 self.gradient_checkpointing = False
74
75 def forward(
76 self,
77 sample: torch.Tensor,
78 image_only_indicator: torch.Tensor,
79 num_frames: int = 1,
80 ) -> torch.Tensor:
81 r"""The forward method of the `Decoder` class."""
82
83 sample = self.conv_in(sample)
84
85 upscale_dtype = next(iter(self.up_blocks.parameters())).dtype
86 # if self.training and self.gradient_checkpointing:
87 if self.gradient_checkpointing:
88
89 def create_custom_forward(module):
90 def custom_forward(*inputs):
91 return module(*inputs)
92
93 return custom_forward
94
95 if is_torch_version(">=", "1.11.0"):
96 # middle
97 sample = torch.utils.checkpoint.checkpoint(
98 create_custom_forward(self.mid_block),
99 sample,
100 image_only_indicator,
101 use_reentrant=False,
102 )
103 sample = sample.to(upscale_dtype)
104
105 # up
106 for up_block in self.up_blocks:
107 sample = torch.utils.checkpoint.checkpoint(
108 create_custom_forward(up_block),
109 sample,
110 image_only_indicator,
111 use_reentrant=False,
112 )
113 else:
114 # middle
115 sample = torch.utils.checkpoint.checkpoint(
116 create_custom_forward(self.mid_block),
117 sample,
118 image_only_indicator,
119 )
120 sample = sample.to(upscale_dtype)
121
122 # up
123 for up_block in self.up_blocks:
124 sample = torch.utils.checkpoint.checkpoint(
125 create_custom_forward(up_block),
126 sample,
127 image_only_indicator,
128 )
129 else:
130 # middle
131 sample = self.mid_block(sample, image_only_indicator=image_only_indicator)
132 sample = sample.to(upscale_dtype)

Callers

nothing calls this directly

Calls 1

checkpointMethod · 0.80

Tested by

no test coverage detected