diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index e3951a777..c70d6fa51 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -816,8 +816,7 @@ class MooncakeKVManager(CommonKVManager): ): # TODO(shangming): Fix me when nvlink_transport of Mooncake is bug-free if ( - self.enable_custom_mem_pool - and self.custom_mem_pool_type in ("NVLINK", "INTRA_NODE_NVLINK") + self.enable_custom_mem_pool and self.custom_mem_pool_type == "NVLINK" ) or envs.SGLANG_MOONCAKE_SEND_AUX_TCP.get(): return self.send_aux_tcp(req, prefill_aux_index, dst_aux_ptrs) diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index d64fd0298..bba4927ed 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -205,7 +205,7 @@ class MetadataBuffers: # TODO(shangming): Fix me (use 'cuda') when nvlink_transport of Mooncake is bug-free device = "cpu" elif envs.SGLANG_MOONCAKE_CUSTOM_MEM_POOL.get() == "INTRA_NODE_NVLINK": - device = "cpu" + device = "cuda" with ( torch.cuda.use_mem_pool(self.custom_mem_pool) if self.custom_mem_pool