diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 64d97f5c6..64e01c9e4 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -876,7 +876,8 @@ 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 == "NVLINK" + self.enable_custom_mem_pool + and self.custom_mem_pool_type in ("NVLINK", "INTRA_NODE_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 d7956a604..6591e743e 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -153,7 +153,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 = "cuda" + device = "cpu" with ( torch.cuda.use_mem_pool(self.custom_mem_pool) if self.custom_mem_pool