[LoRA, Performance] Speedup multi-LoRA serving - Step 1 (#1587)

This commit is contained in:
Ying Sheng
2024-10-06 10:33:44 -07:00
committed by GitHub
parent 58d1082e39
commit 9c064bf78a
3 changed files with 34 additions and 32 deletions
+18 -9
View File
@@ -274,18 +274,24 @@ class LoRAManager:
cur_uids = set(forward_batch.lora_paths)
assert len(cur_uids) <= self.max_loras_per_batch
i = 0
j = len(self.active_uids)
evictable_uids = list(self.active_uids)
for uid in cur_uids:
if uid not in self.active_uids:
while i < len(evictable_uids) and evictable_uids[i] in cur_uids:
i += 1
if i < len(evictable_uids):
if j < self.max_loras_per_batch:
index = j
j += 1
else:
while i < len(evictable_uids) and evictable_uids[i] in cur_uids:
i += 1
assert i < len(evictable_uids)
self.active_uids.remove(evictable_uids[i])
self.buffer_id.pop(evictable_uids[i])
self.load_lora(uid, i)
index = i
i += 1
self.load_lora(uid, index)
self.active_uids.add(uid)
self.buffer_id[uid] = i
i += 1
self.buffer_id[uid] = index
if cur_uids == set([None]):
return
@@ -295,8 +301,11 @@ class LoRAManager:
seg_lens = (
forward_batch.extend_seq_lens
if forward_batch.forward_mode.is_extend()
else torch.ones(bs)
else torch.ones(bs, device="cuda")
)
# FIXME: reuse the data rather than recompute
seg_indptr = torch.zeros((bs + 1,), dtype=torch.int32, device="cuda")
seg_indptr[1:] = torch.cumsum(seg_lens, dim=0)
weight_indices = torch.empty((bs,), dtype=torch.int64, device="cuda")
for i, lora_path in enumerate(forward_batch.lora_paths):
weight_indices[i] = self.buffer_id[lora_path]
@@ -310,7 +319,7 @@ class LoRAManager:
self.A_buffer[weight_name][layer_id],
self.B_buffer[weight_name][layer_id],
bs,
seg_lens,
seg_indptr,
weight_indices,
)
else:
@@ -319,6 +328,6 @@ class LoRAManager:
self.B_buffer["q_proj"][layer_id],
self.B_buffer["kv_proj"][layer_id],
bs,
seg_lens,
seg_indptr,
weight_indices,
)