Add a simple torch native attention backend (#2241)

This commit is contained in:
Qun Yang
2024-12-01 03:01:25 -08:00
committed by GitHub
parent fc78640e00
commit 62c516ac45
7 changed files with 388 additions and 26 deletions
+18 -14
View File
@@ -743,20 +743,24 @@ class ScheduleBatch:
extend_lens = torch.tensor(self.extend_lens, dtype=torch.int32).to(
self.device, non_blocking=True
)
write_req_to_token_pool_triton[(bs,)](
self.req_to_token_pool.req_to_token,
self.req_pool_indices,
pre_lens,
self.seq_lens,
extend_lens,
self.out_cache_loc,
self.req_to_token_pool.req_to_token.shape[1],
)
# The triton kernel is equivalent to the following python code.
# self.req_to_token_pool.write(
# (req.req_pool_idx, slice(pre_len, seq_len)),
# out_cache_loc[pt : pt + req.extend_input_len],
# )
if global_server_args_dict["attention_backend"] != "torch_native":
write_req_to_token_pool_triton[(bs,)](
self.req_to_token_pool.req_to_token,
self.req_pool_indices,
pre_lens,
self.seq_lens,
extend_lens,
self.out_cache_loc,
self.req_to_token_pool.req_to_token.shape[1],
)
else:
pt = 0
for i in range(bs):
self.req_to_token_pool.write(
(self.req_pool_indices[i], slice(pre_lens[i], self.seq_lens[i])),
self.out_cache_loc[pt : pt + self.extend_lens[i]],
)
pt += self.extend_lens[i]
# TODO: some tensors can be reused for ForwardBatchInfo (e.g., extend_lens, cumsum_start)
if self.model_config.is_encoder_decoder: