| 137 | |
| 138 | |
| 139 | def absolute_value_scaling2( |
| 140 | predicted_depth, |
| 141 | ground_truth_depth, |
| 142 | s_init=1.0, |
| 143 | t_init=0.0, |
| 144 | lr=1e-4, |
| 145 | max_iters=1000, |
| 146 | tol=1e-6, |
| 147 | ): |
| 148 | # Initialize s and t as torch tensors with requires_grad=True |
| 149 | s = torch.tensor( |
| 150 | [s_init], |
| 151 | requires_grad=True, |
| 152 | device=predicted_depth.device, |
| 153 | dtype=predicted_depth.dtype, |
| 154 | ) |
| 155 | t = torch.tensor( |
| 156 | [t_init], |
| 157 | requires_grad=True, |
| 158 | device=predicted_depth.device, |
| 159 | dtype=predicted_depth.dtype, |
| 160 | ) |
| 161 | |
| 162 | optimizer = torch.optim.Adam([s, t], lr=lr) |
| 163 | |
| 164 | prev_loss = None |
| 165 | |
| 166 | for i in range(max_iters): |
| 167 | optimizer.zero_grad() |
| 168 | |
| 169 | # Compute predicted aligned depth |
| 170 | predicted_aligned = s * predicted_depth + t |
| 171 | |
| 172 | # Compute absolute error |
| 173 | abs_error = torch.abs(predicted_aligned - ground_truth_depth) |
| 174 | |
| 175 | # Compute loss |
| 176 | loss = torch.sum(abs_error) |
| 177 | |
| 178 | # Backpropagate |
| 179 | loss.backward() |
| 180 | |
| 181 | # Update parameters |
| 182 | optimizer.step() |
| 183 | |
| 184 | # Check convergence |
| 185 | if prev_loss is not None and torch.abs(prev_loss - loss) < tol: |
| 186 | break |
| 187 | |
| 188 | prev_loss = loss.item() |
| 189 | |
| 190 | return s.detach().item(), t.detach().item() |
| 191 | |
| 192 | |
| 193 | def depth_evaluation( |