update lora/dora intermediate var names

This commit is contained in:
tanmayshishodia
2024-07-17 19:44:42 +00:00
parent a59979d3b4
commit 5ae8c7ddb5
3 changed files with 32 additions and 32 deletions
+8 -8
View File
@@ -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)
+5 -5
View File
@@ -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
View File
@@ -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"]