From 5ae8c7ddb5ac8b95b1e78220e36750a3d46d749b Mon Sep 17 00:00:00 2001 From: tanmayshishodia Date: Wed, 17 Jul 2024 19:44:42 +0000 Subject: [PATCH] update lora/dora intermediate var names --- src/adapter/dora_layers.py | 16 ++++++++-------- src/adapter/lora_layers.py | 10 +++++----- src/adapter/utils.py | 38 +++++++++++++++++++------------------- 3 files changed, 32 insertions(+), 32 deletions(-) diff --git a/src/adapter/dora_layers.py b/src/adapter/dora_layers.py index 0573b2e..9a28911 100644 --- a/src/adapter/dora_layers.py +++ b/src/adapter/dora_layers.py @@ -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) diff --git a/src/adapter/lora_layers.py b/src/adapter/lora_layers.py index 1031546..15df5a5 100644 --- a/src/adapter/lora_layers.py +++ b/src/adapter/lora_layers.py @@ -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) diff --git a/src/adapter/utils.py b/src/adapter/utils.py index ed27dd1..ec5ee8e 100644 --- a/src/adapter/utils.py +++ b/src/adapter/utils.py @@ -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"]