[Feature] Support Tensor Parallelism and Weight Slicing for Lora (#4274)
Co-authored-by: ShenAo1111 <1377693092@qq.com> Co-authored-by: Baizhou Zhang <sobereddiezhang@gmail.com>
This commit is contained in:
@@ -782,6 +782,8 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
else:
|
||||
self.num_kv_heads = divide(self.total_num_kv_heads, tp_size)
|
||||
self.num_kv_head_replicas = 1
|
||||
self.q_proj_shard_size = self.num_heads * self.head_size
|
||||
self.kv_proj_shard_size = self.num_kv_heads * self.head_size
|
||||
input_size = self.hidden_size
|
||||
output_size = (
|
||||
(self.num_heads + 2 * self.num_kv_heads) * tp_size * self.head_size
|
||||
|
||||
Reference in New Issue
Block a user