souging/TRELLIS_TextTo3D
0
1import torch2from .z_order import xyz2key as z_order_encode_3from .z_order import key2xyz as z_order_decode_4from .hilbert import encode as hilbert_encode_5from .hilbert import decode as hilbert_decode_6 7 8@torch.inference_mode()9def encode(grid_coord, batch=None, depth=16, order="z"):10 assert order in {"z", "z-trans", "hilbert", "hilbert-trans"}11 if order == "z":12 code = z_order_encode(grid_coord, depth=depth)13 elif order == "z-trans":14 code = z_order_encode(grid_coord[:, [1, 0, 2]], depth=depth)15 elif order == "hilbert":16 code = hilbert_encode(grid_coord, depth=depth)17 elif order == "hilbert-trans":18 code = hilbert_encode(grid_coord[:, [1, 0, 2]], depth=depth)19 else:20 raise NotImplementedError21 if batch is not None:22 batch = batch.long()23 code = batch << depth * 3 | code24 return code25 26 27@torch.inference_mode()28def decode(code, depth=16, order="z"):29 assert order in {"z", "hilbert"}30 batch = code >> depth * 331 code = code & ((1 << depth * 3) - 1)32 if order == "z":33 grid_coord = z_order_decode(code, depth=depth)34 elif order == "hilbert":35 grid_coord = hilbert_decode(code, depth=depth)36 else:37 raise NotImplementedError38 return grid_coord, batch39 40 41def z_order_encode(grid_coord: torch.Tensor, depth: int = 16):42 x, y, z = grid_coord[:, 0].long(), grid_coord[:, 1].long(), grid_coord[:, 2].long()43 # we block the support to batch, maintain batched code in Point class44 code = z_order_encode_(x, y, z, b=None, depth=depth)45 return code46 47 48def z_order_decode(code: torch.Tensor, depth):49 x, y, z, _ = z_order_decode_(code, depth=depth)50 grid_coord = torch.stack([x, y, z], dim=-1) # (N, 3)51 return grid_coord52 53 54def hilbert_encode(grid_coord: torch.Tensor, depth: int = 16):55 return hilbert_encode_(grid_coord, num_dims=3, num_bits=depth)56 57 58def hilbert_decode(code: torch.Tensor, depth: int = 16):59 return hilbert_decode_(code, num_dims=3, num_bits=depth)