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