Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 11 additions & 11 deletions fbgemm_gpu/fbgemm_gpu/batched_unary_embeddings_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,17 +38,17 @@ def __init__(self, num_tasks: int, hash_sizes: list[int], long_index: bool = Fal
# [N][sum(E)][1]
embedding_data = torch.randn(size=(num_tasks, sum(self.hash_sizes), 1))
self.weight = torch.nn.Parameter(embedding_data)
# Keep the offsets as Python ints as well as a buffer. Slicing the
# weight with the buffer's elements would read values off a tensor,
# which yields empty slices on meta / fake tensors (they carry shape
# but no data); the Python copy keeps split_embedding_weights()
# statically shaped so the module is constructible on those devices.
table_offsets = [0]
for hash_size in self.hash_sizes:
table_offsets.append(table_offsets[-1] + hash_size)
self.table_offsets: list[int] = table_offsets
index_dtype = torch.int64 if long_index else torch.int32
table_offsets_tensor = torch.cat(
[
torch.tensor([0], dtype=index_dtype),
torch.cumsum(
torch.tensor(hash_sizes),
dim=0,
dtype=index_dtype,
),
]
)
table_offsets_tensor = torch.tensor(table_offsets, dtype=index_dtype)
self.register_buffer("table_offsets_tensor", table_offsets_tensor)
self.init_parameters()

Expand All @@ -71,7 +71,7 @@ def split_embedding_weights(self):
embedding_weights.append(
self.weight.detach()[
n,
self.table_offsets_tensor[t] : self.table_offsets_tensor[t + 1],
self.table_offsets[t] : self.table_offsets[t + 1],
:,
]
)
Expand Down
46 changes: 46 additions & 0 deletions fbgemm_gpu/test/batched_unary_embeddings_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -239,6 +239,52 @@ def test_gpu(self) -> None:
def test_cpu(self) -> None:
self._test_main(gpu_infer=False)

def test_meta_device_construction(self) -> None:
"""
Constructing on meta must work: split_embedding_weights() slices the
weight per table, and meta tensors carry shape but no data. Slicing
with elements of table_offsets_tensor would read values off a
value-less tensor, collapsing every slice to length 0 and tripping the
shape assert in init_parameters(). The Python-int offsets keep the
slices statically shaped.

Model-analysis tooling constructs modules under torch.device("meta")
to estimate FLOPs and parameter counts without allocating, which is
the path that motivated this test.
"""
hash_sizes = [100, 200]
num_tasks = 3
with torch.device("meta"):
unary_emb = batched_unary_embeddings_ops.BatchedUnaryEmbeddingBag(
num_tasks=num_tasks, hash_sizes=hash_sizes, long_index=True
)

self.assertEqual(unary_emb.weight.shape, (num_tasks, sum(hash_sizes), 1))
split_weights = unary_emb.split_embedding_weights()
self.assertEqual(len(split_weights), num_tasks * len(hash_sizes))
# Order matches init_parameters()'s `hash_sizes * num_tasks` zip.
for i, param in enumerate(split_weights):
self.assertTrue(param.is_meta)
self.assertEqual(param.shape, (hash_sizes[i % len(hash_sizes)], 1))

def test_torchscript_scriptable(self) -> None:
"""
split_embedding_weights() and init_parameters() are @torch.jit.export,
so the module must stay scriptable. Guards the Python-int
table_offsets attribute against TorchScript inference regressions.
"""
hash_sizes = [100, 200]
num_tasks = 3
unary_emb = batched_unary_embeddings_ops.BatchedUnaryEmbeddingBag(
num_tasks=num_tasks, hash_sizes=hash_sizes, long_index=True
)
scripted = torch.jit.script(unary_emb)

split_weights = scripted.split_embedding_weights()
self.assertEqual(len(split_weights), num_tasks * len(hash_sizes))
for i, param in enumerate(split_weights):
self.assertEqual(param.shape, (hash_sizes[i % len(hash_sizes)], 1))

@unittest.skipIf(*gpu_unavailable)
# This test exercises the HIP launch-side limit and requires a large
# output tensor (~17 GiB) plus offsets (~4 GiB) — total ~22 GiB GPU
Expand Down
Loading