change import style

This commit is contained in:
tanmayshishodia
2024-07-16 19:19:47 +00:00
parent 2174a8c69c
commit 39665af7ef
2 changed files with 10 additions and 17 deletions
+5 -8
View File
@@ -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()
+5 -9
View File
@@ -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()