update lora/dora intermediate var names
This commit is contained in:
@@ -37,20 +37,20 @@ class DoraTheta(base_layer.Theta):
|
||||
else:
|
||||
return False
|
||||
|
||||
def _dorafy_var(self, var):
|
||||
def _dorafy_var(self, w):
|
||||
lora_a = super().__getattr__("lora_a")
|
||||
lora_b = super().__getattr__("lora_b")
|
||||
dora_m = super().__getattr__("dora_m")
|
||||
|
||||
new_var = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b)
|
||||
new_var = jnp.reshape(new_var, var.shape)
|
||||
lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b)
|
||||
lora_delta = jnp.reshape(lora_delta, w.shape)
|
||||
|
||||
new_var += var
|
||||
w_prime = w + lora_delta
|
||||
|
||||
column_norm = jnp.linalg.norm(new_var, ord=2, axis=0, keepdims=True)
|
||||
norm_adapted = new_var / column_norm
|
||||
w = dora_m * norm_adapted
|
||||
return w
|
||||
column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True)
|
||||
norm_adapted = w_prime / column_norm
|
||||
w_prime = dora_m * norm_adapted
|
||||
return w_prime
|
||||
|
||||
def __getattr__(self, k):
|
||||
var = super().__getattr__(k)
|
||||
|
||||
@@ -35,13 +35,13 @@ class LoraTheta(base_layer.Theta):
|
||||
else:
|
||||
return False
|
||||
|
||||
def _lorafy_var(self, var):
|
||||
def _lorafy_var(self, w):
|
||||
lora_a = super().__getattr__("lora_a")
|
||||
lora_b = super().__getattr__("lora_b")
|
||||
new_var = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b)
|
||||
new_var = jnp.reshape(new_var, var.shape)
|
||||
new_var += var
|
||||
return new_var
|
||||
lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b)
|
||||
lora_delta = jnp.reshape(lora_delta, w.shape)
|
||||
w_prime = w + lora_delta
|
||||
return w_prime
|
||||
|
||||
def __getattr__(self, k):
|
||||
var = super().__getattr__(k)
|
||||
|
||||
+19
-19
@@ -184,22 +184,22 @@ def _merge_adapter_weights(
|
||||
lora_a = params["lora_a"]
|
||||
lora_b = params["lora_b"]
|
||||
|
||||
var = linear["w"]
|
||||
w = linear["w"]
|
||||
|
||||
new_var = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b)
|
||||
new_var = jnp.reshape(new_var, var.shape)
|
||||
new_var += var
|
||||
lora_delta = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b)
|
||||
lora_delta = jnp.reshape(lora_delta, w.shape)
|
||||
w_prime = w + lora_delta
|
||||
|
||||
if use_dora:
|
||||
dora_m = params["dora_m"]
|
||||
column_norm = jnp.linalg.norm(new_var, ord=2, axis=0, keepdims=True)
|
||||
norm_adapted = new_var / column_norm
|
||||
calc_weights = dora_m * norm_adapted
|
||||
linear["w"] = calc_weights
|
||||
column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True)
|
||||
norm_adapted = w_prime / column_norm
|
||||
w_prime = dora_m * norm_adapted
|
||||
linear["w"] = w_prime
|
||||
del linear["dora_m"]
|
||||
|
||||
else:
|
||||
linear["w"] = new_var
|
||||
linear["w"] = w_prime
|
||||
|
||||
del linear["lora_a"]
|
||||
del linear["lora_b"]
|
||||
@@ -214,22 +214,22 @@ def _merge_adapter_weights(
|
||||
lora_a = params["lora_a"]
|
||||
lora_b = params["lora_b"]
|
||||
|
||||
var = attention[component]["w"]
|
||||
w = attention[component]["w"]
|
||||
|
||||
new_var = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b)
|
||||
new_var = jnp.reshape(new_var, var.shape)
|
||||
new_var += var
|
||||
lora_delta = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b)
|
||||
lora_delta = jnp.reshape(lora_delta, w.shape)
|
||||
w_prime = w + lora_delta
|
||||
|
||||
if use_dora:
|
||||
m = params["dora_m"]
|
||||
column_norm = jnp.linalg.norm(new_var, ord=2, axis=0, keepdims=True)
|
||||
norm_adapted = new_var / column_norm
|
||||
calc_weights = m * norm_adapted
|
||||
attention[component]["w"] = calc_weights
|
||||
dora_m = params["dora_m"]
|
||||
column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True)
|
||||
norm_adapted = w_prime / column_norm
|
||||
w_prime = dora_m * norm_adapted
|
||||
attention[component]["w"] = w_prime
|
||||
del attention[component]["dora_m"]
|
||||
|
||||
else:
|
||||
attention[component]["w"] = new_var
|
||||
attention[component]["w"] = w_prime
|
||||
|
||||
del attention[component]["lora_a"]
|
||||
del attention[component]["lora_b"]
|
||||
|
||||
Reference in New Issue
Block a user