CoolFace
Apppublic

souging/TRELLIS_TextTo3D

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
default.py59 linesDownload Raw Back to pytorch
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)