diff --git a/src/adapter/dora_layers.py b/src/adapter/dora_layers.py index ce35b71..0573b2e 100644 --- a/src/adapter/dora_layers.py +++ b/src/adapter/dora_layers.py @@ -13,14 +13,11 @@ # limitations under the License. from jax import numpy as jnp -from praxis import base_layer, pytypes -from praxis.layers.attentions import AttentionProjection, CombinedQKVProjectionLayer -from praxis.layers.linears import Linear +from praxis import base_layer +from praxis.layers import attentions, linears WeightInit = base_layer.WeightInit -template_field = base_layer.template_field WeightHParams = base_layer.WeightHParams -JTensor = pytypes.JTensor class DoraTheta(base_layer.Theta): @@ -83,7 +80,7 @@ class DoraThetaDescriptor: return DoraTheta(obj) -class DoraLinear(Linear): +class DoraLinear(linears.Linear): rank: int = 0 lora_init: WeightInit | None = None theta = DoraThetaDescriptor() @@ -121,7 +118,7 @@ class DoraLinear(Linear): ) -class DoraAttentionProjection(AttentionProjection): +class DoraAttentionProjection(attentions.AttentionProjection): rank: int = 0 lora_init: WeightInit | None = None theta = DoraThetaDescriptor() @@ -166,7 +163,7 @@ class DoraAttentionProjection(AttentionProjection): ) -class DoraCombinedQKVProjection(CombinedQKVProjectionLayer): +class DoraCombinedQKVProjection(attentions.CombinedQKVProjectionLayer): rank: int = 0 lora_init: WeightInit | None = None theta = DoraThetaDescriptor() diff --git a/src/adapter/lora_layers.py b/src/adapter/lora_layers.py index 7669d06..1031546 100644 --- a/src/adapter/lora_layers.py +++ b/src/adapter/lora_layers.py @@ -13,15 +13,11 @@ # limitations under the License. from jax import numpy as jnp -from praxis import base_layer, pax_fiddle, pytypes -from praxis.layers.attentions import AttentionProjection, CombinedQKVProjectionLayer -from praxis.layers.linears import Linear +from praxis import base_layer +from praxis.layers import attentions, linears WeightInit = base_layer.WeightInit -LayerTpl = pax_fiddle.Config[base_layer.BaseLayer] -template_field = base_layer.template_field WeightHParams = base_layer.WeightHParams -JTensor = pytypes.JTensor class LoraTheta(base_layer.Theta): @@ -75,7 +71,7 @@ class LoraThetaDescriptor: return LoraTheta(obj) -class LoraLinear(Linear): +class LoraLinear(linears.Linear): rank: int = 0 lora_init: WeightInit | None = None theta = LoraThetaDescriptor() @@ -104,7 +100,7 @@ class LoraLinear(Linear): ) -class LoraAttentionProjection(AttentionProjection): +class LoraAttentionProjection(attentions.AttentionProjection): rank: int = 0 lora_init: WeightInit | None = None theta = LoraThetaDescriptor() @@ -140,7 +136,7 @@ class LoraAttentionProjection(AttentionProjection): ) -class LoraCombinedQKVProjection(CombinedQKVProjectionLayer): +class LoraCombinedQKVProjection(attentions.CombinedQKVProjectionLayer): rank: int = 0 lora_init: WeightInit | None = None theta = LoraThetaDescriptor()