INNER CODE UNIT · Python

n_epoch

wenhaochai/StableVideo · app.py:261

            lr, n_epoch = 1e-3, 500
            agg_net = AGGNet().cuda()
            loss_fn = nn.L1Loss()
            optimizer = optim.SGD(agg_net.parameters(), lr=lr, momentum=0.9)
            for _ in range(n_epoch):
                loss = 0.
                for i in range(n_keyframes):
                    e_img = result_list[i]
                    temp_agg_atlas = agg_net(agg_atlas)
                    rec_img = F.grid_sample(temp_agg_atlas[None], 
                                            self.crops['foreground_uvs'][i].reshape(1, -1, 1, 2), 
                                            mode="bilinear", 
                                            align_corners=self.data.config["align_corners"])
                    rec_img = rec_img.clamp(min=0.0, max=1.0).reshape(e_img.shape)
                    loss += loss_fn(rec_img, e_img)
                optimizer.zero_grad()
                loss.backward()
                optimizer.step()

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…