import sys
import torch
from torch.distributed._tensor import distribute_tensor, DTensor, DeviceMesh
from torch.distributed._tensor.placement_types import Replicate, Shard
from torch.testing._internal.common_utils import run_tests, TEST_WITH_DEV_DBG_ASAN
from torch.testing._internal.distributed._tensor.common_dtensor import DTensorTestBase
import torch_npu
from torch_npu.testing.common_distributed import with_comms, skipIfUnsupportMultiNPU
if TEST_WITH_DEV_DBG_ASAN:
raise RuntimeError("Skip dev-asan as torch + multiprocessing spawn have known issues")
class TestEmbeddingOp(DTensorTestBase):
def _run_embedding_op_test(
self,
shard_dim,
input_size,
num_embeddings,
embedding_dim,
**kwargs,
):
device_mesh = DeviceMesh(self.device_type, torch.arange(self.world_size))
torch.manual_seed(0)
local_embedding = torch.nn.Embedding(
num_embeddings,
embedding_dim,
device=self.device_type,
**kwargs,
)
sharded_embedding = torch.nn.Embedding(
num_embeddings,
embedding_dim,
device=self.device_type,
**kwargs,
)
sharded_embedding.weight = torch.nn.Parameter(
distribute_tensor(local_embedding.weight, device_mesh, [Shard(shard_dim)])
)
torch.manual_seed(10)
inp = torch.randint(
0, num_embeddings, tuple(input_size), device=self.device_type
)
target = torch.empty(
*inp.size(), embedding_dim, dtype=torch.float, device=self.device_type
).random_(0, 1)
placements = [Replicate()]
replicate_inp = DTensor.from_local(inp, device_mesh, placements)
sharded_output = sharded_embedding(replicate_inp)
output = sharded_output.redistribute(
sharded_output.device_mesh, [Replicate()]
).to_local()
local_output = local_embedding(inp)
self.assertEqual(local_output, output)
loss = torch.nn.CrossEntropyLoss()
attn_loss = loss(
output,
target,
)
attn_dup_loss = loss(
local_output,
target,
)
attn_loss.backward()
attn_dup_loss.backward()
gradient = sharded_embedding.weight.grad.redistribute(
sharded_output.device_mesh, [Replicate()]
).to_local()
local_grad = local_embedding.weight.grad
self.assertEqual(gradient, local_grad)
local_output = torch.nn.functional.embedding(
inp,
local_embedding.weight,
**kwargs,
)
sharded_output = torch.nn.functional.embedding(
replicate_inp,
sharded_embedding.weight,
**kwargs,
)
self.assertEqual(
local_output,
sharded_output.redistribute(
sharded_output.device_mesh, [Replicate()]
).to_local(),
)
@skipIfUnsupportMultiNPU(4)
@with_comms
def test_sharded_embedding_colwise_errors(self):
with self.assertRaisesRegex(
NotImplementedError,
"DTensor does not support sharded embedding operation with max_norm yet!",
):
self._run_embedding_op_test(
1, [8, 6, 5, 4], 23, 13, padding_idx=12, max_norm=2.0
)
@skipIfUnsupportMultiNPU(4)
@with_comms
def test_sharded_embedding_rowwise(self):
with self.assertRaisesRegex(
NotImplementedError,
"DTensor does not support row-wise sharded embedding operation yet!",
):
self._run_embedding_op_test(0, [5, 12], 16, 22)
if __name__ == "__main__":
run_tests()