Enable overlap scheduler by default for the triton attention backend (#2105)
This commit is contained in:
@@ -170,18 +170,9 @@ class Scheduler:
|
||||
if not self.is_generation:
|
||||
self.enable_overlap = False
|
||||
logger.info("Overlap scheduler is disabled for embedding models.")
|
||||
if (
|
||||
server_args.attention_backend == "triton"
|
||||
or server_args.enable_double_sparsity
|
||||
or (
|
||||
self.model_config.attention_arch == AttentionArch.MLA
|
||||
and not self.server_args.disable_mla
|
||||
)
|
||||
):
|
||||
self.enable_overlap = False
|
||||
logger.info(
|
||||
"Overlap scheduler is disabled if using triton attention backend."
|
||||
)
|
||||
|
||||
if self.enable_overlap:
|
||||
self.disable_jump_forward = True
|
||||
|
||||
# Launch a tensor parallel worker
|
||||
if self.enable_overlap:
|
||||
|
||||
@@ -94,10 +94,21 @@ class TpModelWorkerClient:
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_thread_func_(self):
|
||||
batch_pt = 0
|
||||
batch_lists = [None] * 2
|
||||
|
||||
while True:
|
||||
model_worker_batch, future_token_ids_ct = self.input_queue.get()
|
||||
if not model_worker_batch:
|
||||
break
|
||||
|
||||
# Keep a reference of model_worker_batch by storing it into a list.
|
||||
# Otherwise, the tensor members of model_worker_batch will be released
|
||||
# by pytorch and cause CUDA illegal memory access errors.
|
||||
batch_lists[batch_pt % 2] = model_worker_batch
|
||||
batch_pt += 1
|
||||
|
||||
# Create event
|
||||
self.launch_done = threading.Event()
|
||||
copy_done = torch.cuda.Event()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user