File size: 2,762 Bytes
bc4c433
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
"""Occupancy MLP: raw XYZ → inside/outside logit."""

from __future__ import annotations

import torch
import torch.nn as nn
from torch import Tensor

# Checkpoint schema tag (not a YAML knob).
CHECKPOINT_KIND = "occupancy_mlp"


def build_mlp(in_dim: int, hidden: int, depth: int) -> nn.Sequential:
    """Linear→ReLU × ``depth`` then a 1-logit head. Shared by OccupancyMLP."""
    if in_dim < 1:
        raise ValueError(f"in_dim must be >= 1, got {in_dim}")
    if hidden < 1:
        raise ValueError(f"hidden must be >= 1, got {hidden}")
    if depth < 1:
        raise ValueError(f"depth must be >= 1, got {depth}")
    layers: list[nn.Module] = []
    dim = in_dim
    for _ in range(depth):
        layers.append(nn.Linear(dim, hidden))
        layers.append(nn.ReLU(inplace=True))
        dim = hidden
    layers.append(nn.Linear(dim, 1))
    return nn.Sequential(*layers)


class OccupancyMLP(nn.Module):
    """
    Tiny fully-connected occupancy field.

    Maps a batch of 3D query coordinates to a single unnormalized logit per
    point. A later training step will apply ``binary_cross_entropy_with_logits``
    (do not softmax / sigmoid inside ``forward``).

    Shapes
    ------
    xyz:    ``(B, 3)``  batch of query points (device follows the caller)
    output: ``(B, 1)``  logits; positive → inside, negative → outside

    Device
    ------
    Parameters live on whatever device the module was moved to
    (``.to(device)`` / ``.cuda()``). ``xyz`` must already be on that same
    device; this module does not copy tensors.
    """

    def __init__(self, hidden: int = 64, depth: int = 4) -> None:
        """
        Build Linear→ReLU blocks then a 1-logit head.

        Parameters
        ----------
        hidden:
            Channel width of each hidden Linear (must be ``>= 1``).
        depth:
            Number of hidden Linear+ReLU blocks (must be ``>= 1``).
        """
        super().__init__()
        self.hidden = hidden
        self.depth = depth
        # First Linear is 3 → H; remaining blocks are H → H.
        self.net = build_mlp(3, hidden, depth)

    def forward(self, xyz: Tensor) -> Tensor:
        """
        Evaluate occupancy logits at query coordinates.

        Parameters
        ----------
        xyz:
            Float tensor of shape ``(B, 3)``. Last dim is Cartesian XYZ.

        Returns
        -------
        Tensor
            Float tensor of shape ``(B, 1)`` on the same device as ``xyz``.
        """
        if xyz.ndim != 2 or xyz.shape[-1] != 3:
            raise ValueError(
                f"xyz must have shape (B, 3), got {tuple(xyz.shape)}"
            )
        # Sequential Linear layers require matching dtype/device with parameters.
        return self.net(xyz)