(p, depth_values,axis=1)
| 452 | |
| 453 | |
| 454 | def depth_regression(p, depth_values,axis=1): |
| 455 | if depth_values.dim() <= 2: |
| 456 | # print("regression dim <= 2") |
| 457 | depth_values = depth_values.view(*depth_values.shape, 1, 1) |
| 458 | depth = torch.sum(p * depth_values, axis=axis) |
| 459 | |
| 460 | return depth |
| 461 | |
| 462 | |
| 463 | def winner_take_all(prob_volume, depth_values): |