File size: 3,846 Bytes
76c9728
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
from typing import Literal, Tuple
import math

import torch
import torch.nn as nn
import torch.nn.functional as F

from flex_gemm.nn import (
    SubmanifoldConv3d,
    SparsePool3d,
    SparseUpsample3d,
)


def _pick_spconv_algorithm(in_channels: int, out_channels: int) -> str:
    """Pick the submanifold-conv index-GEMM variant by channel width.

    Wide GEMMs (max(in, out) >= 128) benefit from `masked_implicit_gemm`;
    narrow ones stay on the default `implicit_gemm`.
    """
    if max(in_channels, out_channels) >= 128:
        return "masked_implicit_gemm"
    return "implicit_gemm"


def make_conv3d(in_channels: int, out_channels: int, kernel_size: int = 3) -> SubmanifoldConv3d:
    return SubmanifoldConv3d(
        in_channels, out_channels, kernel_size,
        algorithm=_pick_spconv_algorithm(in_channels, out_channels),
    )


def _with_channels(shape: torch.Size, channels: int) -> torch.Size:
    return torch.Size(list(shape[:-1]) + [channels])


def _ceil_div(value: int, divisor: int) -> int:
    return (value + divisor - 1) // divisor


def _ceil_downsample_shape(shape: torch.Size, sparse_dim: int, factor: int) -> torch.Size:
    sparse_shape = list(shape[:sparse_dim])
    dense_shape = list(shape[sparse_dim:])
    for dim in range(sparse_dim - 3, sparse_dim):
        sparse_shape[dim] = _ceil_div(int(sparse_shape[dim]), factor)
    return torch.Size([*sparse_shape, *dense_shape])


class SparseResBlock3d(nn.Module):
    def __init__(self, channels: int, out_channels: int = None):
        super().__init__()
        self.channels = channels
        self.out_channels = out_channels or channels

        self.norm1 = nn.LayerNorm(channels, elementwise_affine=True, eps=1e-6)
        self.activation_fn = F.silu

        self.conv1 = make_conv3d(channels, self.out_channels)
        self.conv2 = make_conv3d(self.out_channels, self.out_channels)
        self.skip_connection = (
            nn.Linear(channels, self.out_channels)
            if channels != self.out_channels else nn.Identity()
        )

    def init_weights(self):
        for parameter in self.conv2.parameters():
            nn.init.zeros_(parameter)

    def forward(self, feats, coords, shape, neighbor_cache=None):
        h = self.activation_fn(self.norm1(feats).type_as(feats))
        h, neighbor_cache = self.conv1(h, coords, shape, neighbor_cache=neighbor_cache)
        h = self.activation_fn(h)
        h, neighbor_cache = self.conv2(h, coords, shape, neighbor_cache=neighbor_cache)
        return h + self.skip_connection(feats), neighbor_cache


class PoolDown(nn.Module):
    def __init__(self, in_ch: int, out_ch: int, factor: int):
        super().__init__()
        self.factor = factor
        self.pool = SparsePool3d(kernel_size=factor, stride=factor, reduce="mean")
        self.linear = nn.Linear(in_ch, out_ch)

    def forward(self, feats, coords, shape):
        output_shape = _ceil_downsample_shape(shape, coords.shape[1], self.factor)
        feats, coords, shape, down_cache = self.pool(feats, coords, shape, output_shape=output_shape)
        feats = self.linear(feats)
        return feats, coords, _with_channels(shape, feats.shape[-1]), down_cache


class NearestUp(nn.Module):
    def __init__(self, in_ch: int, out_ch: int, factor: int):
        super().__init__()
        self.upsample = SparseUpsample3d(scale_factor=factor, mode="nearest")
        self.linear = nn.Linear(in_ch, out_ch)

    def forward(self, feats, coords, shape, target_coords, target_shape, up_cache=None):
        feats = self.linear(feats)
        shape = _with_channels(shape, feats.shape[-1])
        feats, coords, shape, _ = self.upsample(
            feats, coords, shape,
            output_coords=target_coords, output_shape=target_shape,
            neighbor_cache=up_cache,
        )
        return feats, coords, shape