MCPcopy Create free account
hub / github.com/JasonLSC/GSCodec_Studio / train

Method train

examples/image_fitting.py:77–149  ·  view source on GitHub ↗
(
        self,
        iterations: int = 1000,
        lr: float = 0.01,
        save_imgs: bool = False,
        model_type: Literal["3dgs", "2dgs"] = "3dgs",
    )

Source from the content-addressed store, hash-verified

75 self.viewmat.requires_grad = False
76
77 def train(
78 self,
79 iterations: int = 1000,
80 lr: float = 0.01,
81 save_imgs: bool = False,
82 model_type: Literal["3dgs", "2dgs"] = "3dgs",
83 ):
84 optimizer = optim.Adam(
85 [self.rgbs, self.means, self.scales, self.opacities, self.quats], lr
86 )
87 mse_loss = torch.nn.MSELoss()
88 frames = []
89 times = [0] * 2 # rasterization, backward
90 K = torch.tensor(
91 [
92 [self.focal, 0, self.W / 2],
93 [0, self.focal, self.H / 2],
94 [0, 0, 1],
95 ],
96 device=self.device,
97 )
98
99 if model_type == "3dgs":
100 rasterize_fnc = rasterization
101 elif model_type == "2dgs":
102 rasterize_fnc = rasterization_2dgs
103
104 for iter in range(iterations):
105 start = time.time()
106
107 renders = rasterize_fnc(
108 self.means,
109 self.quats / self.quats.norm(dim=-1, keepdim=True),
110 self.scales,
111 torch.sigmoid(self.opacities),
112 torch.sigmoid(self.rgbs),
113 self.viewmat[None],
114 K[None],
115 self.W,
116 self.H,
117 packed=False,
118 )[0]
119 out_img = renders[0]
120 torch.cuda.synchronize()
121 times[0] += time.time() - start
122 loss = mse_loss(out_img, self.gt_image)
123 optimizer.zero_grad()
124 start = time.time()
125 loss.backward()
126 torch.cuda.synchronize()
127 times[1] += time.time() - start
128 optimizer.step()
129 print(f"Iteration {iter + 1}/{iterations}, Loss: {loss.item()}")
130
131 if save_imgs and iter % 5 == 0:
132 frames.append((out_img.detach().cpu().numpy() * 255).astype(np.uint8))
133 if save_imgs:
134 # save them as a gif with PIL

Callers 1

mainFunction · 0.95

Calls 2

stepMethod · 0.80
backwardMethod · 0.45

Tested by

no test coverage detected