From 229bc65baa522288d0e278990f1ddadb2026b6c7 Mon Sep 17 00:00:00 2001 From: Yuliya Zhautouskaya Date: Thu, 1 Oct 2026 09:24:54 -0700 Subject: [PATCH] Fix Cosmos3 FP8 tensor parallelism --- examples/cosmos3/cosmos_parallel.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/examples/cosmos3/cosmos_parallel.py b/examples/cosmos3/cosmos_parallel.py index 6175005ab029..663dd918388b 100644 --- a/examples/cosmos3/cosmos_parallel.py +++ b/examples/cosmos3/cosmos_parallel.py @@ -431,7 +431,15 @@ def enable_cosmos3_tensor_parallel(transformer, tp_mesh): own processor) on a 2-D ``(tp, cp)`` mesh, or with ``enable_cosmos3_flash_attention`` for TP without CP. """ - from torch.distributed.tensor.parallel import ColwiseParallel, RowwiseParallel, parallelize_module + from torch.distributed.tensor.parallel import ColwiseParallel, parallelize_module + from torch.distributed.tensor.parallel import RowwiseParallel as TorchRowwiseParallel + + # Skip ModelOpt quantizer children to avoid: + # AttributeError: 'TensorQuantizer' object has no attribute 'weight' + class RowwiseParallel(TorchRowwiseParallel): + def _partition_linear_fn(self, name, module, device_mesh): + if isinstance(module, torch.nn.Linear): + return super()._partition_linear_fn(name, module, device_mesh) tp = tp_mesh.size() dev = torch.device("cuda", torch.cuda.current_device())