Skip to content

vllm.distributed.device_communicators.cuda_communicator

Classes:

CudaCommunicator

Bases: DeviceCommunicatorBase

Methods:

  • broadcast –

    Broadcast a tensor from source rank to all ranks.

  • combine –

    Combine the hidden states and router logits from the appropriate device.

  • dispatch –

    Dispatch the hidden states and topk weights/ids to the appropriate device.

  • dispatch_router_logits –

    Dispatch the hidden states and router logits to the appropriate device.

  • recv –

    Receives a tensor from the source rank.

  • send –

    Sends a tensor to the destination rank in a blocking way.

Source code in vllm/distributed/device_communicators/cuda_communicator.py
  35
  36
  37
  38
  39
  40
  41
  42
  43
  44
  45
  46
  47
  48
  49
  50
  51
  52
  53
  54
  55
  56
  57
  58
  59
  60
  61
  62
  63
  64
  65
  66
  67
  68
  69
  70
  71
  72
  73
  74
  75
  76
  77
  78
  79
  80
  81
  82
  83
  84
  85
  86
  87
  88
  89
  90
  91
  92
  93
  94
  95
  96
  97
  98
  99
 100
 101
 102
 103
 104
 105
 106
 107
 108
 109
 110
 111
 112
 113
 114
 115
 116
 117
 118
 119
 120
 121
 122
 123
 124
 125
 126
 127
 128
 129
 130
 131
 132
 133
 134
 135
 136
 137
 138
 139
 140
 141
 142
 143
 144
 145
 146
 147
 148
 149
 150
 151
 152
 153
 154
 155
 156
 157
 158
 159
 160
 161
 162
 163
 164
 165
 166
 167
 168
 169
 170
 171
 172
 173
 174
 175
 176
 177
 178
 179
 180
 181
 182
 183
 184
 185
 186
 187
 188
 189
 190
 191
 192
 193
 194
 195
 196
 197
 198
 199
 200
 201
 202
 203
 204
 205
 206
 207
 208
 209
 210
 211
 212
 213
 214
 215
 216
 217
 218
 219
 220
 221
 222
 223
 224
 225
 226
 227
 228
 229
 230
 231
 232
 233
 234
 235
 236
 237
 238
 239
 240
 241
 242
 243
 244
 245
 246
 247
 248
 249
 250
 251
 252
 253
 254
 255
 256
 257
 258
 259
 260
 261
 262
 263
 264
 265
 266
 267
 268
 269
 270
 271
 272
 273
 274
 275
 276
 277
 278
 279
 280
 281
 282
 283
 284
 285
 286
 287
 288
 289
 290
 291
 292
 293
 294
 295
 296
 297
 298
 299
 300
 301
 302
 303
 304
 305
 306
 307
 308
 309
 310
 311
 312
 313
 314
 315
 316
 317
 318
 319
 320
 321
 322
 323
 324
 325
 326
 327
 328
 329
 330
 331
 332
 333
 334
 335
 336
 337
 338
 339
 340
 341
 342
 343
 344
 345
 346
 347
 348
 349
 350
 351
 352
 353
 354
 355
 356
 357
 358
 359
 360
 361
 362
 363
 364
 365
 366
 367
 368
 369
 370
 371
 372
 373
 374
 375
 376
 377
 378
 379
 380
 381
 382
 383
 384
 385
 386
 387
 388
 389
 390
 391
 392
 393
 394
 395
 396
 397
 398
 399
 400
 401
 402
 403
 404
 405
 406
 407
 408
 409
 410
 411
 412
 413
 414
 415
 416
 417
 418
 419
 420
 421
 422
 423
 424
 425
 426
 427
 428
 429
 430
 431
 432
 433
 434
 435
 436
 437
 438
 439
 440
 441
 442
 443
 444
 445
 446
 447
 448
 449
 450
 451
 452
 453
 454
 455
 456
 457
 458
 459
 460
 461
 462
 463
 464
 465
 466
 467
 468
 469
 470
 471
 472
 473
 474
 475
 476
 477
 478
 479
 480
 481
 482
 483
 484
 485
 486
 487
 488
 489
 490
 491
 492
 493
 494
 495
 496
 497
 498
 499
 500
 501
 502
 503
 504
 505
 506
 507
 508
 509
 510
 511
 512
 513
 514
 515
 516
 517
 518
 519
 520
 521
 522
 523
 524
 525
 526
 527
 528
 529
 530
 531
 532
 533
 534
 535
 536
 537
 538
 539
 540
 541
 542
 543
 544
 545
 546
 547
 548
 549
 550
 551
 552
 553
 554
 555
 556
 557
 558
 559
 560
 561
 562
 563
 564
 565
 566
 567
 568
 569
 570
 571
 572
 573
 574
 575
 576
 577
 578
 579
 580
 581
 582
 583
 584
 585
 586
 587
 588
 589
 590
 591
 592
 593
 594
 595
 596
 597
 598
 599
 600
 601
 602
 603
 604
 605
 606
 607
 608
 609
 610
 611
 612
 613
 614
 615
 616
 617
 618
 619
 620
 621
 622
 623
 624
 625
 626
 627
 628
 629
 630
 631
 632
 633
 634
 635
 636
 637
 638
 639
 640
 641
 642
 643
 644
 645
 646
 647
 648
 649
 650
 651
 652
 653
 654
 655
 656
 657
 658
 659
 660
 661
 662
 663
 664
 665
 666
 667
 668
 669
 670
 671
 672
 673
 674
 675
 676
 677
 678
 679
 680
 681
 682
 683
 684
 685
 686
 687
 688
 689
 690
 691
 692
 693
 694
 695
 696
 697
 698
 699
 700
 701
 702
 703
 704
 705
 706
 707
 708
 709
 710
 711
 712
 713
 714
 715
 716
 717
 718
 719
 720
 721
 722
 723
 724
 725
 726
 727
 728
 729
 730
 731
 732
 733
 734
 735
 736
 737
 738
 739
 740
 741
 742
 743
 744
 745
 746
 747
 748
 749
 750
 751
 752
 753
 754
 755
 756
 757
 758
 759
 760
 761
 762
 763
 764
 765
 766
 767
 768
 769
 770
 771
 772
 773
 774
 775
 776
 777
 778
 779
 780
 781
 782
 783
 784
 785
 786
 787
 788
 789
 790
 791
 792
 793
 794
 795
 796
 797
 798
 799
 800
 801
 802
 803
 804
 805
 806
 807
 808
 809
 810
 811
 812
 813
 814
 815
 816
 817
 818
 819
 820
 821
 822
 823
 824
 825
 826
 827
 828
 829
 830
 831
 832
 833
 834
 835
 836
 837
 838
 839
 840
 841
 842
 843
 844
 845
 846
 847
 848
 849
 850
 851
 852
 853
 854
 855
 856
 857
 858
 859
 860
 861
 862
 863
 864
 865
 866
 867
 868
 869
 870
 871
 872
 873
 874
 875
 876
 877
 878
 879
 880
 881
 882
 883
 884
 885
 886
 887
 888
 889
 890
 891
 892
 893
 894
 895
 896
 897
 898
 899
 900
 901
 902
 903
 904
 905
 906
 907
 908
 909
 910
 911
 912
 913
 914
 915
 916
 917
 918
 919
 920
 921
 922
 923
 924
 925
 926
 927
 928
 929
 930
 931
 932
 933
 934
 935
 936
 937
 938
 939
 940
 941
 942
 943
 944
 945
 946
 947
 948
 949
 950
 951
 952
 953
 954
 955
 956
 957
 958
 959
 960
 961
 962
 963
 964
 965
 966
 967
 968
 969
 970
 971
 972
 973
 974
 975
 976
 977
 978
 979
 980
 981
 982
 983
 984
 985
 986
 987
 988
 989
 990
 991
 992
 993
 994
 995
 996
 997
 998
 999
1000
1001
1002
1003
1004
1005
1006
1007
1008
1009
1010
1011
1012
1013
1014
1015
1016
1017
1018
1019
1020
1021
1022
1023
1024
1025
1026
1027
1028
1029
1030
1031
1032
1033
1034
1035
1036
1037
1038
1039
1040
1041
1042
class CudaCommunicator(DeviceCommunicatorBase):
    def __init__(
        self,
        cpu_group: ProcessGroup,
        device: torch.device | None = None,
        device_group: ProcessGroup | None = None,
        unique_name: str = "",
        global_ranks: list[int] | None = None,
        global_world_size: int | None = None,
        tcp_store_group: StatelessProcessGroup | None = None,
        use_all2all: bool = False,
    ):
        super().__init__(
            cpu_group,
            device,
            device_group,
            unique_name,
            global_ranks,
            global_world_size,
            use_all2all=use_all2all,
        )
        # Match the group name exactly so ETP does not enable TP-only backends.
        if unique_name.split(":")[0] != "tp":
            # custom allreduce or torch symm mem can be used only by tp
            use_custom_allreduce = False
            use_torch_symm_mem = False
            use_flashinfer_allreduce = False
            use_flashinfer_pcie_ipc_allreduce = False
            use_aiter_allreduce = False
        else:
            from vllm.distributed.parallel_state import _ENABLE_CUSTOM_ALL_REDUCE

            use_custom_allreduce = _ENABLE_CUSTOM_ALL_REDUCE
            use_torch_symm_mem = envs.VLLM_ALLREDUCE_USE_SYMM_MEM
            # FlashInfer all-reduce does not provide a fixed reduction order.
            use_flashinfer_allreduce = (
                envs.VLLM_ALLREDUCE_USE_FLASHINFER and not envs.VLLM_BATCH_INVARIANT
            )
            use_flashinfer_pcie_ipc_allreduce = (
                envs.VLLM_ALLREDUCE_USE_FLASHINFER_PCIE_IPC
                and not envs.VLLM_BATCH_INVARIANT
            )
            # Neither AITER nor QuickReduce all-reduce has a fixed reduction order.
            use_aiter_allreduce = (
                use_custom_allreduce
                and not envs.VLLM_BATCH_INVARIANT
                and bool(rocm_aiter_ops.is_custom_all_reduce_enabled())
            )

        self.use_custom_allreduce = use_custom_allreduce
        self.use_torch_symm_mem = use_torch_symm_mem
        self.use_flashinfer_allreduce = use_flashinfer_allreduce
        self.use_flashinfer_pcie_ipc_allreduce = use_flashinfer_pcie_ipc_allreduce
        self.use_aiter_allreduce = use_aiter_allreduce

        # lazy import to avoid documentation build error
        from vllm.distributed.device_communicators.custom_all_reduce import (
            CustomAllreduce,
        )
        from vllm.distributed.device_communicators.flashinfer_all_reduce import (
            FlashInferAllReduce,
        )
        from vllm.distributed.device_communicators.flashinfer_pcie_ipc_all_reduce import (  # noqa: E501
            FlashInferPcieIpcAllReduce,
        )
        from vllm.distributed.device_communicators.pynccl import PyNcclCommunicator
        from vllm.distributed.device_communicators.quick_all_reduce import (
            QuickAllReduce,
        )
        from vllm.distributed.device_communicators.symm_mem import SymmMemCommunicator

        self.pynccl_comm: PyNcclCommunicator | None = None
        if self.world_size > 1:
            self.pynccl_comm = PyNcclCommunicator(
                group=self.cpu_group if tcp_store_group is None else tcp_store_group,
                device=self.device,
            )
            if is_symmetric_memory_enabled():
                register_nccl_symmetric_ops(self.pynccl_comm)

        self.ca_comm: CustomAllreduce | None = None
        self.qr_comm: QuickAllReduce | None = None
        self.symm_mem_comm: SymmMemCommunicator | None = None
        self.fi_ar_comm: FlashInferAllReduce | None = None
        self.fi_pcie_ipc_ar_comm: FlashInferPcieIpcAllReduce | None = None
        self.aiter_ar_comm: AiterCustomAllreduce | None = None
        self.use_aiter_ag_rs: bool = False

        # cuMem graph buffers cannot be IPC-registered; capture copies them instead.
        config = get_current_vllm_config_or_none()
        register = config is None or not config.use_cumem_cudagraph_pool

        if use_torch_symm_mem and current_platform.is_cuda():
            self.symm_mem_comm = SymmMemCommunicator(
                group=self.cpu_group,
                device=self.device,
            )

        if self.use_flashinfer_allreduce and self.world_size > 1:
            self.fi_ar_comm = FlashInferAllReduce(
                group=self.cpu_group,
                device=self.device,
            )

        if (
            self.use_flashinfer_pcie_ipc_allreduce
            and self.world_size > 1
            and self.device_group is not None
        ):
            self.fi_pcie_ipc_ar_comm = FlashInferPcieIpcAllReduce(
                group=self.device_group,
                tune_group=self.cpu_group,
                device=self.device,
            )

        if self.use_aiter_allreduce and self.world_size > 1:
            self.aiter_ar_comm = AiterCustomAllreduce(
                group=self.cpu_group,
                device=self.device,
            )

        if use_custom_allreduce and self.aiter_ar_comm is None and self.world_size > 1:
            # Initialize a custom fast all-reduce implementation.
            self.ca_comm = CustomAllreduce(
                group=self.cpu_group,
                device=self.device,
                symm_mem_enabled=(
                    self.symm_mem_comm is not None and not self.symm_mem_comm.disabled
                ),
                register_graph_buffers=register,
            )

        # AITER custom all-gather/reduce-scatter DP-attention dispatch/combine
        if (
            "dp" in unique_name
            and self.world_size in (2, 4, 8)
            and current_platform.is_rocm()
            and rocm_aiter_ops.is_custom_all_reduce_enabled()
        ):
            self.aiter_ar_comm = AiterCustomAllreduce(
                group=self.cpu_group,
                device=self.device,
            )
            if self.aiter_ar_comm.disabled:
                self.aiter_ar_comm = None
            else:
                self.use_aiter_ag_rs = True

        if (
            use_custom_allreduce
            and self.world_size > 1
            and current_platform.is_rocm()
            and not envs.VLLM_BATCH_INVARIANT
        ):
            # Initialize a custom quick all-reduce implementation for AMD.
            # Quick reduce is designed as a complement to custom allreduce
            # (vLLM's or AITER's), so it is initialized for either backend.
            # Based on quickreduce (https://github.com/mk1-project/quickreduce).
            # On ROCm, 'use_custom_allreduce==True' means it must currently be
            # an MI300 series.
            self.qr_comm = QuickAllReduce(group=self.cpu_group, device=self.device)

        if self.world_size > 1:
            self._log_all_reduce_backend_selection()

        if self.use_all2all:
            if self.all2all_backend in ("naive", "allgather_reducescatter"):
                from .all2all import AgRsAll2AllManager

                self.all2all_manager = AgRsAll2AllManager(
                    self.cpu_group, tcp_store_group
                )
            elif self.all2all_backend == "deepep_high_throughput":
                from .all2all import DeepEPHTAll2AllManager

                self.all2all_manager = DeepEPHTAll2AllManager(
                    self.cpu_group, tcp_store_group
                )
            elif self.all2all_backend == "deepep_low_latency":
                from .all2all import DeepEPLLAll2AllManager

                self.all2all_manager = DeepEPLLAll2AllManager(
                    self.cpu_group, tcp_store_group
                )
            elif self.all2all_backend in (
                "mori_high_throughput",
                "mori_low_latency",
            ):
                from .all2all import MoriAll2AllManager

                self.all2all_manager = MoriAll2AllManager(
                    self.cpu_group, self.all2all_backend
                )
            elif self.all2all_backend == "deepep_v2":
                from .all2all import DeepEPV2All2AllManager

                self.all2all_manager = DeepEPV2All2AllManager(
                    self.cpu_group,
                    tcp_store_group,
                    device_group=self.device_group,
                )
            elif self.all2all_backend == "moonep":
                from .all2all import MoonEPAll2AllManager

                self.all2all_manager = MoonEPAll2AllManager(
                    self.cpu_group,
                    tcp_store_group,
                    device_group=self.device_group,
                )
            elif self.all2all_backend == "nixl_ep":
                from .all2all import NixlEPAll2AllManager

                self.all2all_manager = NixlEPAll2AllManager(
                    self.cpu_group, tcp_store_group
                )
            elif (
                self.all2all_backend == "flashinfer_all2allv"
                or self.all2all_backend == "flashinfer_nvlink_two_sided"
            ):
                if self.all2all_backend == "flashinfer_all2allv":
                    logger.warning_once(
                        "'flashinfer_all2allv' is deprecated and has been renamed to"
                        "'flashinfer_nvlink_two_sided'. It will be removed in a future"
                        "release."
                    )
                from .all2all import FlashInferNVLinkTwoSidedManager

                self.all2all_manager = FlashInferNVLinkTwoSidedManager(
                    self.cpu_group, tcp_store_group
                )
            elif self.all2all_backend == "flashinfer_nvlink_one_sided":
                from .all2all import FlashInferNVLinkOneSidedManager

                self.all2all_manager = FlashInferNVLinkOneSidedManager(self.cpu_group)
            elif self.all2all_backend == "passthrough":
                from .all2all import PassThroughAll2AllManager

                self.all2all_manager = PassThroughAll2AllManager(
                    self.cpu_group, tcp_store_group
                )
            else:
                raise ValueError(f"Unknown all2all backend: {self.all2all_backend}")

            logger.info_once(
                "Using %s all2all manager.",
                self.all2all_manager.__class__.__name__,
                scope="global",
            )

    def _log_all_reduce_backend_selection(self) -> None:
        """Log the all-reduce backends that are active for this group.

        The dispatch chain in ``all_reduce`` tries backends in this order and
        falls through to the next one if the current backend rejects the
        input (size/dtype gates) or is disabled. The list of "enabled"
        backends below is the subset of potential backends that may be
        chosen at dispatch time for this group; the actual per-call choice
        depends on the input tensor.
        """
        all_potential_ar_backends = [
            "FLASHINFER_PCIE_IPC",
            "FLASHINFER",
            "NCCL_SYMM_MEM",
            "QUICK_REDUCE",
            "AITER_CUSTOM",
            "CUSTOM",
            "SYMM_MEM",
            "PYNCCL",
        ]
        enabled_ar_backends: list[str] = []
        if (
            self.fi_pcie_ipc_ar_comm is not None
            and not self.fi_pcie_ipc_ar_comm.disabled
        ):
            enabled_ar_backends.append("FLASHINFER_PCIE_IPC")
        if self.fi_ar_comm is not None and not self.fi_ar_comm.disabled:
            enabled_ar_backends.append("FLASHINFER")
        # Mirror the static preconditions of `should_nccl_symm_mem_allreduce`:
        # VLLM_BATCH_INVARIANT off, NCCL symm mem enabled, world_size meets
        # min_world_size, and world_size either has a tuned entry in
        # `custom_ar_preferred_ranges` or is greater than
        # `always_use_above_world_size`. World sizes that fail the latter (e.g.
        # 5/6/7 with the default config) never dispatch NCCL symm mem
        # regardless of input. The per-tensor-size check inside the function
        # stays as a runtime decision.
        nccl_symm_ws_ok = self.world_size >= NCCL_SYMM_MEM_ALL_REDUCE_CONFIG[
            "min_world_size"
        ] and (
            self.world_size
            in NCCL_SYMM_MEM_ALL_REDUCE_CONFIG["custom_ar_preferred_ranges"]
            or self.world_size
            > NCCL_SYMM_MEM_ALL_REDUCE_CONFIG["always_use_above_world_size"]
        )
        if (
            self.pynccl_comm is not None
            and not self.pynccl_comm.disabled
            and is_symmetric_memory_enabled()
            and not envs.VLLM_BATCH_INVARIANT
            and nccl_symm_ws_ok
        ):
            enabled_ar_backends.append("NCCL_SYMM_MEM")
        if self.qr_comm is not None and not self.qr_comm.disabled:
            enabled_ar_backends.append("QUICK_REDUCE")
        if (
            self.use_aiter_allreduce
            and self.aiter_ar_comm is not None
            and not self.aiter_ar_comm.disabled
        ):
            enabled_ar_backends.append("AITER_CUSTOM")
        if self.ca_comm is not None and not self.ca_comm.disabled:
            enabled_ar_backends.append("CUSTOM")
        if self.symm_mem_comm is not None and not self.symm_mem_comm.disabled:
            enabled_ar_backends.append("SYMM_MEM")
        if self.pynccl_comm is not None and not self.pynccl_comm.disabled:
            enabled_ar_backends.append("PYNCCL")

        logger.info_once(
            "Using %s all-reduce backends (in dispatch order) for group "
            "'%s' out of potential backends: %s.",
            "[" + ", ".join(f"'{b}'" for b in enabled_ar_backends) + "]",
            self.unique_name or "<unnamed>",
            "[" + ", ".join(f"'{b}'" for b in all_potential_ar_backends) + "]",
            scope="global",
        )

    def all_reduce(self, input_):
        fi_ar_comm = self.fi_ar_comm
        use_fi_ar = (
            fi_ar_comm is not None
            and not fi_ar_comm.disabled
            and fi_ar_comm.should_use_fi_ar(input_)
        )

        # since currently we perform copy input -> symm_input -> out-of-place AR
        # return symm_output, we don't need to check if input is symmetric
        if (
            self.pynccl_comm is not None
            and not use_fi_ar
            and should_nccl_symm_mem_allreduce(self.pynccl_comm.world_size, input_)
        ):
            out = torch.ops.vllm.all_reduce_symmetric_with_copy(input_)
            if out is not None:
                return out
        qr_comm = self.qr_comm
        if (
            qr_comm is not None
            and not qr_comm.disabled
            and qr_comm.should_quick_allreduce(input_)
        ):
            out = qr_comm.quick_all_reduce(input_)
            assert out is not None
            return out
        fi_pcie_ipc_ar_comm = self.fi_pcie_ipc_ar_comm
        if fi_pcie_ipc_ar_comm is not None and fi_pcie_ipc_ar_comm.should_use(input_):
            return fi_pcie_ipc_ar_comm.all_reduce(input_)
        if use_fi_ar:
            assert fi_ar_comm is not None
            out = fi_ar_comm.all_reduce(input_)
            assert out is not None
            return out
        aiter_ar_comm = self.aiter_ar_comm
        if (
            self.use_aiter_allreduce
            and aiter_ar_comm is not None
            and not aiter_ar_comm.disabled
            and aiter_ar_comm.should_custom_ar(input_)
        ):
            out = aiter_ar_comm.custom_all_reduce(input_)
            assert out is not None
            return out
        ca_comm = self.ca_comm
        if (
            ca_comm is not None
            and not ca_comm.disabled
            and ca_comm.should_custom_ar(input_)
        ):
            out = ca_comm.custom_all_reduce(input_)
            assert out is not None
            return out
        symm_mem_comm = self.symm_mem_comm
        if symm_mem_comm is not None and symm_mem_comm.should_use_symm_mem(input_):
            out = symm_mem_comm.all_reduce(input_)
            assert out is not None
            return out
        pynccl_comm = self.pynccl_comm
        if pynccl_comm is None or pynccl_comm.disabled:
            out = input_.clone()
            torch.distributed.all_reduce(out, group=self.device_group)
            return out
        assert pynccl_comm is not None
        out = pynccl_comm.all_reduce(input_)
        if out is None:
            # fall back to the default all-reduce using PyTorch.
            # this usually happens during testing.
            # when we run the model, allreduce only happens for the TP
            # group, where we always have either custom allreduce or pynccl.
            out = input_.clone()
            torch.distributed.all_reduce(out, group=self.device_group)
        return out

    def custom_all_gather(self, input_: torch.Tensor) -> torch.Tensor | None:
        ca_comm = self.ca_comm
        if ca_comm is None:
            return None
        return ca_comm.custom_all_gather(input_.contiguous())

    def custom_reduce_scatter(self, input_: torch.Tensor) -> torch.Tensor | None:
        ca_comm = self.ca_comm
        if ca_comm is None:
            return None
        return ca_comm.custom_reduce_scatter(input_.contiguous())

    def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
        # Route uniform dim-0 all-gathers through NVLS symmetric memory when
        # enabled (mirrors reduce_scatter); otherwise fall back to the
        # PyNccl/base-class all-gather. Sequence parallelism's
        # gather-before-GEMM uses dim=0 with tp-aligned (uniform) shards.
        if dim < 0:
            dim += input_.dim()
        if dim == 0 and should_nccl_symm_mem_ag_rs():
            return self._all_gather_symm_mem(input_.contiguous())

        pynccl_comm = self.pynccl_comm
        if pynccl_comm is None or pynccl_comm.disabled:
            return super().all_gather(input_, dim)

        # On ROCm, the base-class all_gather (all_gather_into_tensor) is faster
        # than the manual pynccl + torch.empty + movedim + reshape path below,
        # which adds a per-call output allocation and (for dim != 0) an extra
        # copy on every step. This is on the hot path for TP forward passes, so
        # keep ROCm on the base-class collective to avoid a decode regression.
        if current_platform.is_rocm():
            return super().all_gather(input_, dim)

        input_size = input_.size()
        output_size = (input_size[0] * self.world_size,) + input_size[1:]
        output_tensor = torch.empty(
            output_size, dtype=input_.dtype, device=input_.device
        )
        pynccl_comm.all_gather(output_tensor, input_.contiguous())
        output_tensor = output_tensor.reshape((self.world_size,) + input_size)
        output_tensor = output_tensor.movedim(0, dim)
        return output_tensor.reshape(
            input_size[:dim]
            + (self.world_size * input_size[dim],)
            + input_size[dim + 1 :]
        )

    def reduce_scatter(self, input_: torch.Tensor, dim: int = -1):
        world_size = self.world_size
        pynccl_comm = self.pynccl_comm
        assert pynccl_comm is not None
        if dim < 0:
            # Convert negative dim to positive.
            dim += input_.dim()

        # Note: This will produce an incorrect answer if we don't make
        # the input_tensor contiguous. Possible bug in reduce_scatter_tensor?
        input_tensor = input_.movedim(0, dim).contiguous()

        assert input_tensor.shape[0] % world_size == 0
        chunk_size = input_tensor.shape[0] // world_size
        output_shape = (chunk_size,) + input_tensor.shape[1:]

        if should_nccl_symm_mem_ag_rs():
            output = self._reduce_scatter_symm_mem(input_tensor)
        else:
            output = torch.empty(
                output_shape, dtype=input_tensor.dtype, device=input_tensor.device
            )
            pynccl_comm.reduce_scatter(output, input_tensor)

        # Reshape before returning
        return output.movedim(0, dim).contiguous()

    def reduce_scatterv(
        self, input_: torch.Tensor, dim: int = -1, sizes: list[int] | None = None
    ):
        return self._reduce_scatterv(input_, dim, sizes)

    def reduce_scatterv_into_output(
        self,
        input_: torch.Tensor,
        output: torch.Tensor,
        dim: int = -1,
        sizes: list[int] | None = None,
    ) -> torch.Tensor:
        return self._reduce_scatterv(input_, dim, sizes, output)

    def _reduce_scatterv(
        self,
        input_: torch.Tensor,
        dim: int,
        sizes: list[int] | None,
        output: torch.Tensor | None = None,
    ) -> torch.Tensor:
        world_size = self.world_size
        pynccl_comm = self.pynccl_comm
        assert pynccl_comm is not None
        assert not pynccl_comm.disabled, "reduce_scatterv requires PyNccl"
        if dim < 0:
            # Convert negative dim to positive.
            dim += input_.dim()

        sizes_are_explicit = sizes is not None
        # 'sizes' is not needed if all inputs in the same group have the same
        # shape
        if sizes is not None and all(s == sizes[0] for s in sizes):
            sizes = None

        # Note: This will produce an incorrect answer if we don't make
        # the input_tensor contiguous. Possible bug in reduce_scatter_tensor?
        input_tensor = input_.movedim(0, dim).contiguous()

        if sizes is not None:
            assert len(sizes) == world_size, f"{len(sizes)} == {world_size}"
            assert input_tensor.shape[0] == sum(sizes)
            chunk_size = sizes[self.rank_in_group]
        else:
            assert input_tensor.shape[0] % world_size == 0
            chunk_size = input_tensor.shape[0] // world_size
        output_shape = (chunk_size,) + input_tensor.shape[1:]
        if output is not None:
            assert dim == 0, "preallocated reduce-scatter output requires dim=0"
            assert output.shape == output_shape
            assert output.dtype == input_tensor.dtype
            assert output.device == input_tensor.device
            assert output.is_contiguous()

        if self._can_use_aiter_ag_rs(sizes):
            aiter_comm = self.aiter_ar_comm
            assert aiter_comm is not None
            if aiter_comm.should_custom_rs(input_tensor, dim=0):
                if output is None:
                    output = torch.empty(
                        output_shape,
                        dtype=input_tensor.dtype,
                        device=input_tensor.device,
                    )
                aiter_comm.custom_reduce_scatter(input_tensor, output, dim=0)
                return output.movedim(0, dim).contiguous()

        # Symmetric memory is only used when all ranks have uniform sizes.
        # ncclCommWindowRegister is collective: asymmetric pool allocations
        # from variable per-rank sizes cause deadlocks.
        uniform_sizes = sizes is None or all(size == sizes[0] for size in sizes)
        direct_output_supported = (
            output is None
            or is_symmetric_memory_tensor(output)
            or pynccl_comm.nccl_version >= NCCL_DIRECT_SYMM_RS_OUTPUT_MIN_VERSION
        )
        use_symm_mem = (
            uniform_sizes
            and direct_output_supported
            and (
                output is None
                or not sizes_are_explicit
                or is_symmetric_memory_tensor(input_tensor)
            )
            and should_nccl_symm_mem_ag_rs()
        )
        if use_symm_mem:
            output = self._reduce_scatter_symm_mem(input_tensor, output)
        else:
            if output is None:
                output = torch.empty(
                    output_shape, dtype=input_tensor.dtype, device=input_tensor.device
                )
            use_deterministic_rs = envs.VLLM_BATCH_INVARIANT and world_size > 2
            if use_deterministic_rs:
                # Reduce to a fixed root (0) for determinism
                reduced = torch.empty_like(input_tensor)
                scatter_sizes = sizes
                if scatter_sizes is None:
                    scatter_sizes = [chunk_size] * world_size
                pynccl_comm.reduce(reduced, input_tensor, root=0)
                pynccl_comm.scatter(output, reduced, scatter_sizes, root=0)
            elif not uniform_sizes:
                assert sizes is not None
                pynccl_comm.reduce_scatterv(output, input_tensor, sizes=sizes)
            else:
                pynccl_comm.reduce_scatter(output, input_tensor)

        # Reshape before returning
        return output.movedim(0, dim).contiguous()

    def _get_symm_scratch(
        self,
        role: str,
        shape: tuple[int, ...],
        dtype: torch.dtype,
        device: torch.device,
    ) -> torch.Tensor:
        """Persistent, pre-registered NCCL symmetric-memory scratch buffer.

        Allocating a fresh symm tensor per collective pays the
        ``nccl_symm_mem_context`` snapshot + window-registration scan on every
        call (~0.5 ms/RS+AG pair, dwarfing the NVLS transfer itself). Instead,
        each ``(role, shape[1:], dtype, device)`` keeps a geometric high-water
        allocation and returns a leading slice. Superseded allocations remain
        referenced because a captured CUDA graph may still use their pointers;
        geometric growth bounds their total capacity to less than twice the
        current allocation.

        Safe across serial MoE layers and eager sequence parallelism: the
        producer and collective are ordered on the same stream before the next
        same-role operation reuses the buffer. DBO microbatches use distinct
        cache entries. Any future cross-layer communication overlap must also
        use distinct roles.
        """
        from vllm.distributed.device_communicators.pynccl_allocator import (
            nccl_symm_mem_context,
        )
        from vllm.v1.worker.ubatching import dbo_current_ubatch_id

        pynccl_comm = self.pynccl_comm
        assert pynccl_comm is not None
        assert shape, "symmetric scratch buffers require at least one dimension"
        cache = self.__dict__.setdefault("_symm_scratch_bufs", {})
        key = (role, dbo_current_ubatch_id(), tuple(shape[1:]), dtype, device)
        buf = cache.get(key)
        requested_rows = shape[0]
        if buf is None or buf.shape[0] < requested_rows:
            capacity = (
                requested_rows if buf is None else max(requested_rows, 2 * buf.shape[0])
            )
            with nccl_symm_mem_context(pynccl_comm):
                new_buf = torch.empty(
                    (capacity, *shape[1:]), dtype=dtype, device=device
                )
            if buf is not None:
                retired = self.__dict__.setdefault("_retired_symm_scratch_bufs", {})
                retired.setdefault(key, []).append(buf)
            buf = new_buf
            cache[key] = buf
        return buf[:requested_rows]

    def _reduce_scatter_symm_mem(
        self,
        input_tensor: torch.Tensor,
        output: torch.Tensor | None = None,
    ) -> torch.Tensor:
        """ReduceScatter using NCCL symmetric memory (NVLS).

        The MoE reduce_scatterv path passes explicit uniform sizes and calls
        this only with an already-registered input, so that path never stages.
        Uniform calls without a caller-owned output retain their existing
        opt-in behavior and stage ordinary input into registered scratch.
        The output may be ordinary memory.
        """
        pynccl_comm = self.pynccl_comm
        assert pynccl_comm is not None
        chunk = input_tensor.shape[0] // self.world_size
        output_shape = (chunk,) + tuple(input_tensor.shape[1:])

        if is_symmetric_memory_tensor(input_tensor):
            symm_input = input_tensor
        else:
            symm_input = self._get_symm_scratch(
                "rs_in",
                tuple(input_tensor.shape),
                input_tensor.dtype,
                input_tensor.device,
            )
            symm_input.copy_(input_tensor)

        if output is None:
            output = self._get_symm_scratch(
                "rs_out", output_shape, input_tensor.dtype, input_tensor.device
            )
        pynccl_comm.reduce_scatter(output, symm_input)
        return output

    def get_symmetric_memory_buffer(
        self,
        role: str,
        shape: tuple[int, ...],
        dtype: torch.dtype,
        device: torch.device,
    ) -> torch.Tensor | None:
        pynccl_comm = self.pynccl_comm
        if (
            pynccl_comm is None
            or pynccl_comm.disabled
            or pynccl_comm.world_size == 1
            or pynccl_comm.nccl_version < NCCL_DIRECT_SYMM_RS_OUTPUT_MIN_VERSION
            or not should_nccl_symm_mem_ag_rs()
        ):
            return None
        output = self._get_symm_scratch(role, shape, dtype, device)
        return output if is_symmetric_memory_tensor(output) else None

    def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
        """Sends a tensor to the destination rank in a blocking way."""
        """NOTE: `dst` is the local rank of the destination rank."""
        if dst is None:
            dst = (self.rank_in_group + 1) % self.world_size

        pynccl_comm = self.pynccl_comm
        if pynccl_comm is not None and not pynccl_comm.disabled:
            pynccl_comm.send(tensor, dst)
        else:
            torch.distributed.send(tensor, self.ranks[dst], self.device_group)

    def recv(
        self, size: torch.Size, dtype: torch.dtype, src: int | None = None
    ) -> torch.Tensor:
        """Receives a tensor from the source rank."""
        """NOTE: `src` is the local rank of the source rank."""
        if src is None:
            src = (self.rank_in_group - 1) % self.world_size

        tensor = torch.empty(size, dtype=dtype, device=self.device)
        pynccl_comm = self.pynccl_comm
        if pynccl_comm is not None and not pynccl_comm.disabled:
            pynccl_comm.recv(tensor, src)
        else:
            torch.distributed.recv(tensor, self.ranks[src], self.device_group)
        return tensor

    def broadcast(self, tensor: torch.Tensor, src: int = 0) -> torch.Tensor:
        """Broadcast a tensor from source rank to all ranks."""
        if self.world_size == 1:
            return tensor

        pynccl_comm = self.pynccl_comm
        if pynccl_comm is not None and not pynccl_comm.disabled:
            pynccl_comm.broadcast(tensor, src)
            return tensor
        else:
            raise ValueError("No PyNCCL communicator found")

    def destroy(self):
        if self.pynccl_comm is not None:
            self.pynccl_comm.destroy()
            self.pynccl_comm = None
        if self.ca_comm is not None:
            self.ca_comm = None
        if self.aiter_ar_comm is not None:
            self.aiter_ar_comm.close()
            self.aiter_ar_comm = None
        if self.fi_ar_comm is not None:
            self.fi_ar_comm.destroy()
            self.fi_ar_comm = None
        if self.fi_pcie_ipc_ar_comm is not None:
            self.fi_pcie_ipc_ar_comm.destroy()
            self.fi_pcie_ipc_ar_comm = None
        if self.all2all_manager is not None:
            self.all2all_manager.destroy()
            self.all2all_manager = None  # type: ignore[assignment]

    def _can_use_aiter_ag_rs(self, sizes: list[int] | None) -> bool:
        """Whether the AITER custom AG/RS fast path may run for this collective.

        Requires:
        - uniform batches
        - FULL CUDAgraphs
        """
        if (
            not self.use_aiter_ag_rs
            or self.aiter_ar_comm is None
            or self.aiter_ar_comm.disabled
        ):
            return False
        if sizes is not None:
            return False

        from vllm.config.compilation import CUDAGraphMode
        from vllm.forward_context import get_forward_context

        try:
            ctx = get_forward_context()
        except AssertionError:
            return False
        if ctx.cudagraph_runtime_mode != CUDAGraphMode.FULL:
            return False
        bd = ctx.batch_descriptor
        return bd is not None and bd.uniform

    def suspend(self) -> None:
        if self.pynccl_comm is not None:
            self.pynccl_comm.suspend()

    def resume(self) -> None:
        if self.pynccl_comm is not None:
            self.pynccl_comm.resume()

    def checkpoint_prepare(self) -> None:
        # Only FlashInfer all-reduce and FlashInfer all2all are supported for now.
        from .flashinfer_all_reduce import checkpoint_prepare_fi_ar_workspaces

        checkpoint_prepare_fi_ar_workspaces(self.cpu_group)
        if self.all2all_manager is not None:
            self.all2all_manager.checkpoint_prepare()

    def checkpoint_restore(self) -> None:
        # Only FlashInfer all-reduce and FlashInfer all2all are supported for now.
        from .flashinfer_all_reduce import checkpoint_restore_fi_ar_workspaces

        checkpoint_restore_fi_ar_workspaces(self.cpu_group)
        if self.all2all_manager is not None:
            self.all2all_manager.checkpoint_restore()

    def all_gatherv(
        self,
        input_: torch.Tensor | list[torch.Tensor],
        dim: int = 0,
        sizes: list[int] | None = None,
    ):
        if dim != 0:
            raise NotImplementedError("only dim 0 all-gatherv is supported")
        world_size = self.world_size

        # 'sizes' is not needed if all inputs in the same group have the same
        # shape
        if sizes is not None and all(s == sizes[0] for s in sizes):
            sizes = None

        if self._can_use_aiter_ag_rs(sizes):
            aiter_comm = self.aiter_ar_comm
            assert aiter_comm is not None
            if isinstance(input_, torch.Tensor):
                if aiter_comm.should_custom_ag(input_):
                    out = aiter_comm.custom_all_gather(input_, dim=0)
                    if out is not None:
                        return out
            elif all(aiter_comm.should_custom_ag(inp) for inp in input_):
                outs = [aiter_comm.custom_all_gather(inp, dim=0) for inp in input_]
                if all(o is not None for o in outs):
                    return outs

        pynccl_comm = self.pynccl_comm
        assert pynccl_comm is not None and not pynccl_comm.disabled

        # Symmetric memory is only used when all ranks have uniform sizes.
        # ncclCommWindowRegister is collective: asymmetric pool allocations
        # from variable per-rank sizes cause deadlocks.
        if sizes is None and should_nccl_symm_mem_ag_rs():
            if isinstance(input_, torch.Tensor):
                return self._all_gather_symm_mem(input_)
            return self._all_gather_batched_symm_mem(input_)

        def _all_gather_single(input_: torch.Tensor, sizes: list[int] | None = None):
            input_size = input_.size()
            if sizes is not None:
                assert len(sizes) == world_size
                assert input_.shape[dim] == sizes[self.rank_in_group], (
                    f"{input_.shape[dim]} != {sizes[self.rank_in_group]}"
                )
                output_size = (sum(sizes),) + input_size[1:]
            else:
                output_size = (input_size[0] * world_size,) + input_size[1:]
            # Allocate output tensor.
            output_tensor = torch.empty(
                output_size, dtype=input_.dtype, device=input_.device
            )
            if sizes is not None:
                pynccl_comm.all_gatherv(output_tensor, input_, sizes=sizes)
            else:
                pynccl_comm.all_gather(output_tensor, input_)
            return output_tensor

        if isinstance(input_, torch.Tensor):
            return _all_gather_single(input_, sizes)

        output_list = []
        pynccl_comm.group_start()
        for inp in input_:
            output_list.append(_all_gather_single(inp, sizes=sizes))
        pynccl_comm.group_end()

        return output_list

    def _all_gather_symm_mem(self, input_: torch.Tensor) -> torch.Tensor:
        """AllGather a single tensor using NCCL symmetric memory (NVLS).

        Only the output needs to be in symmetric memory; NCCL does not
        require the AG input to be symmetrically allocated.
        """
        pynccl_comm = self.pynccl_comm
        assert pynccl_comm is not None

        out_size = (input_.size(0) * self.world_size,) + tuple(input_.size()[1:])
        # Persistent pre-registered scratch avoids the per-call symm-mem context
        # snapshot/registration overhead (see _get_symm_scratch).
        symm_output = self._get_symm_scratch(
            "ag_out", out_size, input_.dtype, input_.device
        )
        pynccl_comm.all_gather(symm_output, input_)
        return symm_output

    def _all_gather_batched_symm_mem(
        self, inputs: list[torch.Tensor]
    ) -> list[torch.Tensor]:
        """AllGather a list of tensors using NCCL symmetric memory (NVLS).

        Uses group_start/group_end to batch the collectives.
        Only the output needs to be in symmetric memory (see
        _all_gather_symm_mem).
        """
        from vllm.distributed.device_communicators.pynccl_allocator import (
            nccl_symm_mem_context,
        )

        pynccl_comm = self.pynccl_comm
        assert pynccl_comm is not None
        world_size = self.world_size

        symm_outputs = []
        with nccl_symm_mem_context(pynccl_comm):
            for inp in inputs:
                out_size = (inp.size(0) * world_size,) + inp.size()[1:]
                symm_outputs.append(
                    torch.empty(out_size, dtype=inp.dtype, device=inp.device)
                )

        pynccl_comm.group_start()
        for symm_out, inp in zip(symm_outputs, inputs):
            pynccl_comm.all_gather(symm_out, inp)
        pynccl_comm.group_end()

        return symm_outputs

    def dispatch_router_logits(
        self,
        hidden_states: torch.Tensor,
        router_logits: torch.Tensor,
        is_sequence_parallel: bool = False,
        extra_tensors: list[torch.Tensor] | None = None,
    ) -> (
        tuple[torch.Tensor, torch.Tensor]
        | tuple[torch.Tensor, torch.Tensor, list[torch.Tensor]]
    ):
        """Dispatch the hidden states and router logits to the appropriate device.
        This is a no-op in the base class.
        """
        assert self.all2all_manager is not None
        return self.all2all_manager.dispatch_router_logits(
            hidden_states,
            router_logits,
            is_sequence_parallel,
            extra_tensors,
        )

    def dispatch(
        self,
        hidden_states: torch.Tensor,
        topk_weights: torch.Tensor,
        topk_ids: torch.Tensor,
        is_sequence_parallel: bool = False,
        extra_tensors: list[torch.Tensor] | None = None,
    ) -> (
        tuple[torch.Tensor, torch.Tensor, torch.Tensor]
        | tuple[torch.Tensor, torch.Tensor, torch.Tensor, list[torch.Tensor]]
    ):
        """Dispatch the hidden states and topk weights/ids to the appropriate device.
        This is a no-op in the base class.
        """
        assert self.all2all_manager is not None
        return self.all2all_manager.dispatch(
            hidden_states,
            topk_weights,
            topk_ids,
            is_sequence_parallel,
            extra_tensors=extra_tensors,
        )

    def combine(
        self, hidden_states: torch.Tensor, is_sequence_parallel: bool = False
    ) -> torch.Tensor:
        """Combine the hidden states and router logits from the appropriate device.
        This is a no-op in the base class.
        """
        assert self.all2all_manager is not None
        return self.all2all_manager.combine(
            hidden_states,
            is_sequence_parallel,
        )

    def allocate_combine_input(
        self,
        shape: tuple[int, ...],
        dtype: torch.dtype,
        device: torch.device,
        is_sequence_parallel: bool = False,
    ) -> torch.Tensor | None:
        assert self.all2all_manager is not None
        return self.all2all_manager.allocate_combine_input(
            shape, dtype, device, is_sequence_parallel
        )

    def combine_into_output(
        self,
        hidden_states: torch.Tensor,
        output: torch.Tensor,
        is_sequence_parallel: bool = False,
    ) -> torch.Tensor:
        assert self.all2all_manager is not None
        return self.all2all_manager.combine_into_output(
            hidden_states, output, is_sequence_parallel
        )

    def batch_isend_irecv(self, p2p_ops: list):
        pynccl_comm = self.pynccl_comm
        if pynccl_comm is not None and not pynccl_comm.disabled:
            pynccl_comm.batch_isend_irecv(p2p_ops)
        else:
            raise ValueError("No PyNCCL communicator found")

_all_gather_batched_symm_mem(inputs)

AllGather a list of tensors using NCCL symmetric memory (NVLS).

Uses group_start/group_end to batch the collectives. Only the output needs to be in symmetric memory (see _all_gather_symm_mem).

Source code in vllm/distributed/device_communicators/cuda_communicator.py
def _all_gather_batched_symm_mem(
    self, inputs: list[torch.Tensor]
) -> list[torch.Tensor]:
    """AllGather a list of tensors using NCCL symmetric memory (NVLS).

    Uses group_start/group_end to batch the collectives.
    Only the output needs to be in symmetric memory (see
    _all_gather_symm_mem).
    """
    from vllm.distributed.device_communicators.pynccl_allocator import (
        nccl_symm_mem_context,
    )

    pynccl_comm = self.pynccl_comm
    assert pynccl_comm is not None
    world_size = self.world_size

    symm_outputs = []
    with nccl_symm_mem_context(pynccl_comm):
        for inp in inputs:
            out_size = (inp.size(0) * world_size,) + inp.size()[1:]
            symm_outputs.append(
                torch.empty(out_size, dtype=inp.dtype, device=inp.device)
            )

    pynccl_comm.group_start()
    for symm_out, inp in zip(symm_outputs, inputs):
        pynccl_comm.all_gather(symm_out, inp)
    pynccl_comm.group_end()

    return symm_outputs

_all_gather_symm_mem(input_)

AllGather a single tensor using NCCL symmetric memory (NVLS).

Only the output needs to be in symmetric memory; NCCL does not require the AG input to be symmetrically allocated.

Source code in vllm/distributed/device_communicators/cuda_communicator.py
def _all_gather_symm_mem(self, input_: torch.Tensor) -> torch.Tensor:
    """AllGather a single tensor using NCCL symmetric memory (NVLS).

    Only the output needs to be in symmetric memory; NCCL does not
    require the AG input to be symmetrically allocated.
    """
    pynccl_comm = self.pynccl_comm
    assert pynccl_comm is not None

    out_size = (input_.size(0) * self.world_size,) + tuple(input_.size()[1:])
    # Persistent pre-registered scratch avoids the per-call symm-mem context
    # snapshot/registration overhead (see _get_symm_scratch).
    symm_output = self._get_symm_scratch(
        "ag_out", out_size, input_.dtype, input_.device
    )
    pynccl_comm.all_gather(symm_output, input_)
    return symm_output

_can_use_aiter_ag_rs(sizes)

Whether the AITER custom AG/RS fast path may run for this collective.

Requires: - uniform batches - FULL CUDAgraphs

Source code in vllm/distributed/device_communicators/cuda_communicator.py
def _can_use_aiter_ag_rs(self, sizes: list[int] | None) -> bool:
    """Whether the AITER custom AG/RS fast path may run for this collective.

    Requires:
    - uniform batches
    - FULL CUDAgraphs
    """
    if (
        not self.use_aiter_ag_rs
        or self.aiter_ar_comm is None
        or self.aiter_ar_comm.disabled
    ):
        return False
    if sizes is not None:
        return False

    from vllm.config.compilation import CUDAGraphMode
    from vllm.forward_context import get_forward_context

    try:
        ctx = get_forward_context()
    except AssertionError:
        return False
    if ctx.cudagraph_runtime_mode != CUDAGraphMode.FULL:
        return False
    bd = ctx.batch_descriptor
    return bd is not None and bd.uniform

_get_symm_scratch(role, shape, dtype, device)

Persistent, pre-registered NCCL symmetric-memory scratch buffer.

Allocating a fresh symm tensor per collective pays the nccl_symm_mem_context snapshot + window-registration scan on every call (~0.5 ms/RS+AG pair, dwarfing the NVLS transfer itself). Instead, each (role, shape[1:], dtype, device) keeps a geometric high-water allocation and returns a leading slice. Superseded allocations remain referenced because a captured CUDA graph may still use their pointers; geometric growth bounds their total capacity to less than twice the current allocation.

Safe across serial MoE layers and eager sequence parallelism: the producer and collective are ordered on the same stream before the next same-role operation reuses the buffer. DBO microbatches use distinct cache entries. Any future cross-layer communication overlap must also use distinct roles.

Source code in vllm/distributed/device_communicators/cuda_communicator.py
def _get_symm_scratch(
    self,
    role: str,
    shape: tuple[int, ...],
    dtype: torch.dtype,
    device: torch.device,
) -> torch.Tensor:
    """Persistent, pre-registered NCCL symmetric-memory scratch buffer.

    Allocating a fresh symm tensor per collective pays the
    ``nccl_symm_mem_context`` snapshot + window-registration scan on every
    call (~0.5 ms/RS+AG pair, dwarfing the NVLS transfer itself). Instead,
    each ``(role, shape[1:], dtype, device)`` keeps a geometric high-water
    allocation and returns a leading slice. Superseded allocations remain
    referenced because a captured CUDA graph may still use their pointers;
    geometric growth bounds their total capacity to less than twice the
    current allocation.

    Safe across serial MoE layers and eager sequence parallelism: the
    producer and collective are ordered on the same stream before the next
    same-role operation reuses the buffer. DBO microbatches use distinct
    cache entries. Any future cross-layer communication overlap must also
    use distinct roles.
    """
    from vllm.distributed.device_communicators.pynccl_allocator import (
        nccl_symm_mem_context,
    )
    from vllm.v1.worker.ubatching import dbo_current_ubatch_id

    pynccl_comm = self.pynccl_comm
    assert pynccl_comm is not None
    assert shape, "symmetric scratch buffers require at least one dimension"
    cache = self.__dict__.setdefault("_symm_scratch_bufs", {})
    key = (role, dbo_current_ubatch_id(), tuple(shape[1:]), dtype, device)
    buf = cache.get(key)
    requested_rows = shape[0]
    if buf is None or buf.shape[0] < requested_rows:
        capacity = (
            requested_rows if buf is None else max(requested_rows, 2 * buf.shape[0])
        )
        with nccl_symm_mem_context(pynccl_comm):
            new_buf = torch.empty(
                (capacity, *shape[1:]), dtype=dtype, device=device
            )
        if buf is not None:
            retired = self.__dict__.setdefault("_retired_symm_scratch_bufs", {})
            retired.setdefault(key, []).append(buf)
        buf = new_buf
        cache[key] = buf
    return buf[:requested_rows]

_log_all_reduce_backend_selection()

Log the all-reduce backends that are active for this group.

The dispatch chain in all_reduce tries backends in this order and falls through to the next one if the current backend rejects the input (size/dtype gates) or is disabled. The list of "enabled" backends below is the subset of potential backends that may be chosen at dispatch time for this group; the actual per-call choice depends on the input tensor.

Source code in vllm/distributed/device_communicators/cuda_communicator.py
def _log_all_reduce_backend_selection(self) -> None:
    """Log the all-reduce backends that are active for this group.

    The dispatch chain in ``all_reduce`` tries backends in this order and
    falls through to the next one if the current backend rejects the
    input (size/dtype gates) or is disabled. The list of "enabled"
    backends below is the subset of potential backends that may be
    chosen at dispatch time for this group; the actual per-call choice
    depends on the input tensor.
    """
    all_potential_ar_backends = [
        "FLASHINFER_PCIE_IPC",
        "FLASHINFER",
        "NCCL_SYMM_MEM",
        "QUICK_REDUCE",
        "AITER_CUSTOM",
        "CUSTOM",
        "SYMM_MEM",
        "PYNCCL",
    ]
    enabled_ar_backends: list[str] = []
    if (
        self.fi_pcie_ipc_ar_comm is not None
        and not self.fi_pcie_ipc_ar_comm.disabled
    ):
        enabled_ar_backends.append("FLASHINFER_PCIE_IPC")
    if self.fi_ar_comm is not None and not self.fi_ar_comm.disabled:
        enabled_ar_backends.append("FLASHINFER")
    # Mirror the static preconditions of `should_nccl_symm_mem_allreduce`:
    # VLLM_BATCH_INVARIANT off, NCCL symm mem enabled, world_size meets
    # min_world_size, and world_size either has a tuned entry in
    # `custom_ar_preferred_ranges` or is greater than
    # `always_use_above_world_size`. World sizes that fail the latter (e.g.
    # 5/6/7 with the default config) never dispatch NCCL symm mem
    # regardless of input. The per-tensor-size check inside the function
    # stays as a runtime decision.
    nccl_symm_ws_ok = self.world_size >= NCCL_SYMM_MEM_ALL_REDUCE_CONFIG[
        "min_world_size"
    ] and (
        self.world_size
        in NCCL_SYMM_MEM_ALL_REDUCE_CONFIG["custom_ar_preferred_ranges"]
        or self.world_size
        > NCCL_SYMM_MEM_ALL_REDUCE_CONFIG["always_use_above_world_size"]
    )
    if (
        self.pynccl_comm is not None
        and not self.pynccl_comm.disabled
        and is_symmetric_memory_enabled()
        and not envs.VLLM_BATCH_INVARIANT
        and nccl_symm_ws_ok
    ):
        enabled_ar_backends.append("NCCL_SYMM_MEM")
    if self.qr_comm is not None and not self.qr_comm.disabled:
        enabled_ar_backends.append("QUICK_REDUCE")
    if (
        self.use_aiter_allreduce
        and self.aiter_ar_comm is not None
        and not self.aiter_ar_comm.disabled
    ):
        enabled_ar_backends.append("AITER_CUSTOM")
    if self.ca_comm is not None and not self.ca_comm.disabled:
        enabled_ar_backends.append("CUSTOM")
    if self.symm_mem_comm is not None and not self.symm_mem_comm.disabled:
        enabled_ar_backends.append("SYMM_MEM")
    if self.pynccl_comm is not None and not self.pynccl_comm.disabled:
        enabled_ar_backends.append("PYNCCL")

    logger.info_once(
        "Using %s all-reduce backends (in dispatch order) for group "
        "'%s' out of potential backends: %s.",
        "[" + ", ".join(f"'{b}'" for b in enabled_ar_backends) + "]",
        self.unique_name or "<unnamed>",
        "[" + ", ".join(f"'{b}'" for b in all_potential_ar_backends) + "]",
        scope="global",
    )

_reduce_scatter_symm_mem(input_tensor, output=None)

ReduceScatter using NCCL symmetric memory (NVLS).

The MoE reduce_scatterv path passes explicit uniform sizes and calls this only with an already-registered input, so that path never stages. Uniform calls without a caller-owned output retain their existing opt-in behavior and stage ordinary input into registered scratch. The output may be ordinary memory.

Source code in vllm/distributed/device_communicators/cuda_communicator.py
def _reduce_scatter_symm_mem(
    self,
    input_tensor: torch.Tensor,
    output: torch.Tensor | None = None,
) -> torch.Tensor:
    """ReduceScatter using NCCL symmetric memory (NVLS).

    The MoE reduce_scatterv path passes explicit uniform sizes and calls
    this only with an already-registered input, so that path never stages.
    Uniform calls without a caller-owned output retain their existing
    opt-in behavior and stage ordinary input into registered scratch.
    The output may be ordinary memory.
    """
    pynccl_comm = self.pynccl_comm
    assert pynccl_comm is not None
    chunk = input_tensor.shape[0] // self.world_size
    output_shape = (chunk,) + tuple(input_tensor.shape[1:])

    if is_symmetric_memory_tensor(input_tensor):
        symm_input = input_tensor
    else:
        symm_input = self._get_symm_scratch(
            "rs_in",
            tuple(input_tensor.shape),
            input_tensor.dtype,
            input_tensor.device,
        )
        symm_input.copy_(input_tensor)

    if output is None:
        output = self._get_symm_scratch(
            "rs_out", output_shape, input_tensor.dtype, input_tensor.device
        )
    pynccl_comm.reduce_scatter(output, symm_input)
    return output

broadcast(tensor, src=0)

Broadcast a tensor from source rank to all ranks.

Source code in vllm/distributed/device_communicators/cuda_communicator.py
def broadcast(self, tensor: torch.Tensor, src: int = 0) -> torch.Tensor:
    """Broadcast a tensor from source rank to all ranks."""
    if self.world_size == 1:
        return tensor

    pynccl_comm = self.pynccl_comm
    if pynccl_comm is not None and not pynccl_comm.disabled:
        pynccl_comm.broadcast(tensor, src)
        return tensor
    else:
        raise ValueError("No PyNCCL communicator found")

combine(hidden_states, is_sequence_parallel=False)

Combine the hidden states and router logits from the appropriate device. This is a no-op in the base class.

Source code in vllm/distributed/device_communicators/cuda_communicator.py
def combine(
    self, hidden_states: torch.Tensor, is_sequence_parallel: bool = False
) -> torch.Tensor:
    """Combine the hidden states and router logits from the appropriate device.
    This is a no-op in the base class.
    """
    assert self.all2all_manager is not None
    return self.all2all_manager.combine(
        hidden_states,
        is_sequence_parallel,
    )

dispatch(hidden_states, topk_weights, topk_ids, is_sequence_parallel=False, extra_tensors=None)

Dispatch the hidden states and topk weights/ids to the appropriate device. This is a no-op in the base class.

Source code in vllm/distributed/device_communicators/cuda_communicator.py
def dispatch(
    self,
    hidden_states: torch.Tensor,
    topk_weights: torch.Tensor,
    topk_ids: torch.Tensor,
    is_sequence_parallel: bool = False,
    extra_tensors: list[torch.Tensor] | None = None,
) -> (
    tuple[torch.Tensor, torch.Tensor, torch.Tensor]
    | tuple[torch.Tensor, torch.Tensor, torch.Tensor, list[torch.Tensor]]
):
    """Dispatch the hidden states and topk weights/ids to the appropriate device.
    This is a no-op in the base class.
    """
    assert self.all2all_manager is not None
    return self.all2all_manager.dispatch(
        hidden_states,
        topk_weights,
        topk_ids,
        is_sequence_parallel,
        extra_tensors=extra_tensors,
    )

dispatch_router_logits(hidden_states, router_logits, is_sequence_parallel=False, extra_tensors=None)

Dispatch the hidden states and router logits to the appropriate device. This is a no-op in the base class.

Source code in vllm/distributed/device_communicators/cuda_communicator.py
def dispatch_router_logits(
    self,
    hidden_states: torch.Tensor,
    router_logits: torch.Tensor,
    is_sequence_parallel: bool = False,
    extra_tensors: list[torch.Tensor] | None = None,
) -> (
    tuple[torch.Tensor, torch.Tensor]
    | tuple[torch.Tensor, torch.Tensor, list[torch.Tensor]]
):
    """Dispatch the hidden states and router logits to the appropriate device.
    This is a no-op in the base class.
    """
    assert self.all2all_manager is not None
    return self.all2all_manager.dispatch_router_logits(
        hidden_states,
        router_logits,
        is_sequence_parallel,
        extra_tensors,
    )

recv(size, dtype, src=None)

Receives a tensor from the source rank.

Source code in vllm/distributed/device_communicators/cuda_communicator.py
def recv(
    self, size: torch.Size, dtype: torch.dtype, src: int | None = None
) -> torch.Tensor:
    """Receives a tensor from the source rank."""
    """NOTE: `src` is the local rank of the source rank."""
    if src is None:
        src = (self.rank_in_group - 1) % self.world_size

    tensor = torch.empty(size, dtype=dtype, device=self.device)
    pynccl_comm = self.pynccl_comm
    if pynccl_comm is not None and not pynccl_comm.disabled:
        pynccl_comm.recv(tensor, src)
    else:
        torch.distributed.recv(tensor, self.ranks[src], self.device_group)
    return tensor

send(tensor, dst=None)

Sends a tensor to the destination rank in a blocking way.

Source code in vllm/distributed/device_communicators/cuda_communicator.py
def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
    """Sends a tensor to the destination rank in a blocking way."""
    """NOTE: `dst` is the local rank of the destination rank."""
    if dst is None:
        dst = (self.rank_in_group + 1) % self.world_size

    pynccl_comm = self.pynccl_comm
    if pynccl_comm is not None and not pynccl_comm.disabled:
        pynccl_comm.send(tensor, dst)
    else:
        torch.distributed.send(tensor, self.ranks[dst], self.device_group)