pytorch / pytorch/pytorch

DTensor `F.conv2d` fails with `bias=None` - AssertionError in sharding propagation

Open
#167,090 2 comments 0 reactions 0 assignees View on GitHub
bot-triaged module: dtensor oncall: distributed oncall: distributed parallelisms ptd-bot-triaged triaged
Dominant language
Python
Stars
103k
Forks
29.5k
PR merge metrics
PR metrics pending

Description

### 🐛 Describe the bug

`torch.nn.functional.conv2d` fails when using DTensor inputs with `bias=None`, raising an assertion error:
```
AssertionError: assert isinstance(bias_spec, DTensorSpec)
```
The error occurs in the sharding propagation logic for convolution operations, which assumes that `bias` is always a DTensor rather than allowing `None` for bias-free convolutions.

```python
import torch
from torch.distributed.tensor import DTensor, Replicate, distribute_tensor, init_device_mesh
import torch.nn.functional as F

# construct a device mesh with available devices (multi-host or single host)
device_mesh = init_device_mesh("cuda", (2, 1))
placement = [Replicate(), Replicate()]

# With square kernels and equal stride
inputs = torch.randn(1, 4, 5, 5)
kernel = torch.randn(8, 4, 3, 3)
bias = torch.randn(8)

d_inputs = distribute_tensor(inputs, device_mesh=device_mesh, placements=placement)
d_kernel = distribute_tensor(kernel, device_mesh=device_mesh, placements=placement)
d_bias = distribute_tensor(bias, device_mesh=device_mesh, placements=placement)

F.conv2d(input=d_inputs, weight=d_kernel, bias=d_bias, padding=1) # works
F.conv2d(input=inputs, weight=kernel, bias=None, padding=1) # works
F.conv2d(input=d_inputs, weight=d_kernel, bias=None, padding=1) # fails
```
### Steps to Run

```bash
python3 -m torch.distributed.run --nproc_per_node=2
```

### Full Error Stack Trace

```
[rank0]: Traceback (most recent call last):
[rank0]: File ".venv/lib/python3.12/site-packages/torch/distributed/tensor/_sharding_prop.py", line 517, in propagate_op_sharding_non_cached
[rank0]: output_sharding = sharding_prop_func(op_schema)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File ".venv/lib/python3.12/site-packages/torch/distributed/tensor/_ops/_conv_ops.py", line 29, in convolution_rules
[rank0]: assert isinstance(bias_spec, DTensorSpec)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: AssertionError
[rank0]:
[rank0]: The above exception was the direct cause of the following exception:
[rank0]:
[rank0]: Traceback (most recent call last):
[rank0]: File "conv2d_biasless_example.py", line 20, in
[rank0]: F.conv2d(input=d_inputs, weight=d_kernel, bias=None, padding=1)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File ".venv/lib/python3.12/site-packages/torch/distributed/tensor/_api.py", line 349, in __torch_dispatch__
[rank0]: return DTensor._op_dispatcher.dispatch(
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File ".venv/lib/python3.12/site-packages/torch/distributed/tensor/_dispatch.py", line 149, in dispatch
[rank0]: return self._custom_op_handlers[op_call](op_call, args, kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File ".venv/lib/python3.12/site-packages/torch/distributed/tensor/_tp_conv.py", line 239, in convolution_handler
[rank0]: dtensor.DTensor._op_dispatcher.sharding_propagator.propagate(op_info)
[rank0]: File ".venv/lib/python3.12/site-packages/torch/distributed/tensor/_sharding_prop.py", line 327, in propagate
[rank0]: OutputSharding, self.propagate_op_sharding(op_info.schema)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File ".venv/lib/python3.12/site-packages/torch/distributed/tensor/_sharding_prop.py", line 46, in __call__
[rank0]: return self.cache(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File ".venv/lib/python3.12/site-packages/torch/distributed/tensor/_sharding_prop.py", line 521, in propagate_op_sharding_non_cached
[rank0]: raise RuntimeError(
[rank0]: RuntimeError: Sharding propagation failed on op Op(op=aten.convolution.default,
[rank0]: args_schema=Spec((Replicate(), Replicate()) on (1, 4, 5, 5)),
[rank0]: Spec((Replicate(), Replicate()) on (8, 4, 3, 3)),
[rank0]: None, [1, 1], [1, 1], [1, 1], False, [0, 0], 1 @ mesh: (2, 1)).
[rank0]: Error:

### Versions

Collecting environment information...
PyTorch version: 2.9.0+cu128
Is debug build: False
CUDA used to build PyTorch: 12.8
ROCM used to build PyTorch: N/A

OS: Ubuntu 22.04.5 LTS (x86_64)
GCC version: (Ubuntu 13.1.0-8ubuntu1~22.04) 13.1.0
Clang version: Could not collect
CMake version: Could not collect
Libc version: glibc-2.35

Python version: 3.12.12 (main, Oct 14 2025, 21:25:31) [Clang 20.1.4 ] (64-bit runtime)
Python platform: Linux-6.8.0-1043-gcp-x86_64-with-glibc2.35
Is CUDA available: True
CUDA runtime version: Could not collect
CUDA_MODULE_LOADING set to:
GPU models and configuration:
GPU 0: NVIDIA A100-SXM4-80GB
GPU 1: NVIDIA A100-SXM4-80GB

Nvidia driver version: 570.195.03
cuDNN version: Could not collect
Is XPU available: False
HIP runtime version: N/A
MIOpen runtime version: N/A
Is XNNPACK available: True

CPU:
Architecture: x86_64
CPU op-mode(s): 32-bit, 64-bit
Address sizes: 46 bits physical, 48 bits virtual
Byte Order: Little Endian
CPU(s): 24
On-line CPU(s) list: 0-23
Vendor ID: GenuineIntel
Model name: Intel(R) Xeon(R) CPU @ 2.20GHz
CPU family: 6
Model: 85
Thread(s) per core: 2
Core(s) per socket: 12
Socket(s): 1
Stepping: 7
BogoMIPS: 4400.30
Flags: fpu vme de pse tsc msr pae mce cx8 apic sep mtrr pge mca cmov pat pse36 clflush mmx fxsr sse sse2 ss ht syscall nx pdpe1gb rdtscp lm constant_tsc rep_good nopl xtopology nonstop_tsc cpuid tsc_known_freq pni pclmulqdq ssse3 fma cx16 pcid sse4_1 sse4_2 x2apic movbe popcnt aes xsave avx f16c rdrand hypervisor lahf_lm abm 3dnowprefetch ssbd ibrs ibpb stibp ibrs_enhanced fsgsbase tsc_adjust bmi1 hle avx2 smep bmi2 erms invpcid rtm mpx avx512f avx512dq rdseed adx smap clflushopt clwb avx512cd avx512bw avx512vl xsaveopt xsavec xgetbv1 xsaves arat avx512_vnni md_clear arch_capabilities
Hypervisor vendor: KVM
Virtualization type: full
L1d cache: 384 KiB (12 instances)
L1i cache: 384 KiB (12 instances)
L2 cache: 12 MiB (12 instances)
L3 cache: 38.5 MiB (1 instance)
NUMA node(s): 1
NUMA node0 CPU(s): 0-23
Vulnerability Gather data sampling: Not affected
Vulnerability Itlb multihit: Not affected
Vulnerability L1tf: Not affected
Vulnerability Mds: Not affected
Vulnerability Meltdown: Not affected
Vulnerability Mmio stale data: Vulnerable: Clear CPU buffers attempted, no microcode; SMT Host state unknown
Vulnerability Reg file data sampling: Not affected
Vulnerability Retbleed: Mitigation; Enhanced IBRS
Vulnerability Spec rstack overflow: Not affected
Vulnerability Spec store bypass: Mitigation; Speculative Store Bypass disabled via prctl
Vulnerability Spectre v1: Mitigation; usercopy/swapgs barriers and __user pointer sanitization
Vulnerability Spectre v2: Mitigation; Enhanced / Automatic IBRS; IBPB conditional; RSB filling; PBRSB-eIBRS SW sequence; BHI SW loop, KVM SW loop
Vulnerability Srbds: Not affected
Vulnerability Tsx async abort: Vulnerable: Clear CPU buffers attempted, no microcode; SMT Host state unknown
Vulnerability Vmscape: Not affected

Versions of relevant libraries:
[pip3] Could not collect
[conda] Could not collect

cc @awgu @wanchaol @fegin @fduwjj @wz337 @wconstab @d4l3k @pragupta @msaroufim @dcci @aditvenk @weifengpy @tianyu-l @XilunWu @SherlockNoMad @ppwwyyxx @H-Huang

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.