MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / LoraMixin

Class LoraMixin

SwissArmyTransformer/sat/model/finetune/lora2.py:181–255  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

179 return new_lin.cuda() if torch.cuda.is_available() else new_lin
180
181class 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():

Callers 2

__init__Method · 0.90
__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected