| 179 | return new_lin.cuda() if torch.cuda.is_available() else new_lin |
| 180 | |
| 181 | class LoraMixin(BaseMixin): |
| 182 | def __init__(self, |
| 183 | layer_num, |
| 184 | r: int = 0, |
| 185 | lora_alpha: int = 1, |
| 186 | lora_dropout: float = 0., |
| 187 | layer_range = None, |
| 188 | qlora = False, |
| 189 | cross_attention = True): |
| 190 | super().__init__() |
| 191 | self.r = r |
| 192 | self.lora_alpha = lora_alpha |
| 193 | self.lora_dropout = lora_dropout |
| 194 | |
| 195 | if layer_range is None: |
| 196 | layer_range = [i for i in range(layer_num)] |
| 197 | self.layer_range = layer_range |
| 198 | |
| 199 | self.scaling = self.lora_alpha / self.r |
| 200 | self.qlora = qlora |
| 201 | self.cross_attention = cross_attention |
| 202 | |
| 203 | def reinit(self, parent_model): |
| 204 | for i in self.layer_range: |
| 205 | print_rank0(f'replacing layer {i} attention with lora') |
| 206 | parent_model.transformer.layers[i].attention.dense = replace_linear_with_lora(parent_model.transformer.layers[i].attention.dense, 1, self.r, self.lora_alpha, self.lora_dropout, qlora=self.qlora, in_size=parent_model.transformer.hidden_size, out_size=None) |
| 207 | parent_model.transformer.layers[i].attention.query_key_value = replace_linear_with_lora(parent_model.transformer.layers[i].attention.query_key_value, parent_model.transformer.layers[i].attention.stride, self.r, self.lora_alpha, self.lora_dropout, qlora=self.qlora, in_size=parent_model.transformer.hidden_size, out_size=None if not parent_model.transformer.num_multi_query_heads else parent_model.transformer.layers[i].attention.inner_hidden_size + parent_model.transformer.layers[i].attention.hidden_size_per_attention_head * parent_model.transformer.layers[i].attention.num_multi_query_heads * 2) |
| 208 | if self.cross_attention and parent_model.transformer.layers[i].is_decoder: |
| 209 | print_rank0(f'replacing layer {i} cross attention with lora') |
| 210 | kv_size = parent_model.transformer.layers[i].cross_attention.inner_hidden_size * 2 if not parent_model.transformer.cross_num_multi_query_heads else parent_model.transformer.layers[i].cross_attention.hidden_size_per_attention_head * parent_model.transformer.layers[i].cross_attention.cross_num_multi_query_heads * 2 |
| 211 | parent_model.transformer.layers[i].cross_attention.dense = replace_linear_with_lora(parent_model.transformer.layers[i].cross_attention.dense, 1, self.r, self.lora_alpha, self.lora_dropout, qlora=self.qlora, in_size=parent_model.transformer.layers[i].cross_attention.inner_hidden_size, out_size=parent_model.transformer.hidden_size) |
| 212 | parent_model.transformer.layers[i].cross_attention.query = replace_linear_with_lora(parent_model.transformer.layers[i].cross_attention.query, 1, self.r, self.lora_alpha, self.lora_dropout, qlora=self.qlora, in_size=parent_model.transformer.hidden_size, out_size=parent_model.transformer.layers[i].cross_attention.inner_hidden_size) |
| 213 | parent_model.transformer.layers[i].cross_attention.key_value = replace_linear_with_lora(parent_model.transformer.layers[i].cross_attention.key_value, 2, self.r, self.lora_alpha, self.lora_dropout, qlora=self.qlora, in_size=parent_model.transformer.layers[i].cross_attention.cross_attn_hidden_size, out_size=kv_size) |
| 214 | if self.qlora: |
| 215 | print_rank0('replacing chatglm linear layer with 4bit') |
| 216 | def replace_linear_with_nf4(model, name=None, cache={}): |
| 217 | if type(model) in (nn.Linear, RowParallelLinear, ColumnParallelLinear): |
| 218 | out_dim, in_dim = model.weight.shape |
| 219 | bias = model.bias is not None |
| 220 | new_linear = HackLinearNF4(in_dim, out_dim, bias=bias) |
| 221 | new_linear.weight.data.copy_(model.weight.data.detach().clone()) |
| 222 | if bias: |
| 223 | new_linear.bias.data.copy_(model.bias.data.detach().clone()) |
| 224 | return new_linear |
| 225 | names = set() |
| 226 | for name, child in model.named_children(): |
| 227 | if name not in names: |
| 228 | if child in cache: |
| 229 | new_child = cache[child] |
| 230 | else: |
| 231 | new_child = replace_linear_with_nf4(child, name=name, cache=cache) |
| 232 | cache[child] = new_child |
| 233 | setattr(model, name, new_child) |
| 234 | names.add(name) |
| 235 | flag = True |
| 236 | while flag: |
| 237 | flag = False |
| 238 | for name, child in model.named_children(): |