MCPcopy Create free account
hub / github.com/NVlabs/RADIO / __init__

Method __init__

radio/feature_normalizer.py:50–59  ·  view source on GitHub ↗
(self, num_intermediates: int, embed_dim: int, rot_per_layer: bool = False, dtype: torch.dtype = torch.float32)

Source from the content-addressed store, hash-verified

48
49class IntermediateFeatureNormalizer(IntermediateFeatureNormalizerBase):
50 def __init__(self, num_intermediates: int, embed_dim: int, rot_per_layer: bool = False, dtype: torch.dtype = torch.float32):
51 super().__init__()
52 self.register_buffer('alphas', torch.ones(num_intermediates, dtype=dtype))
53
54 rot = torch.eye(embed_dim, dtype=dtype)
55 if rot_per_layer:
56 rot = rot.unsqueeze(0).repeat(num_intermediates, 1, 1)
57
58 self.register_buffer('rotation', rot.contiguous())
59 self.register_buffer('means', torch.zeros(num_intermediates, embed_dim, dtype=dtype))
60
61 def forward(self, x: torch.Tensor, index: int, rot_index: int = None, skip: Optional[int] = None) -> InterFeatState:
62 if rot_index is None:

Callers

nothing calls this directly

Calls 1

__init__Method · 0.45

Tested by

no test coverage detected