mirror of
https://github.com/p-e-w/heretic.git
synced 2026-09-29 07:21:27 -07:00
feat(ara): add optional row-norm preservation
This commit is contained in:
@@ -205,7 +205,7 @@ class Settings(BaseSettings):
|
|||||||
)
|
)
|
||||||
|
|
||||||
use_piqa: bool = Field(
|
use_piqa: bool = Field(
|
||||||
default=True,
|
default=False,
|
||||||
description=(
|
description=(
|
||||||
"Whether to use the Physical Interaction: Question Answering (PIQA) benchmark "
|
"Whether to use the Physical Interaction: Question Answering (PIQA) benchmark "
|
||||||
"as the quality metric instead of the Kullback-Leibler divergence."
|
"as the quality metric instead of the Kullback-Leibler divergence."
|
||||||
@@ -221,7 +221,7 @@ class Settings(BaseSettings):
|
|||||||
)
|
)
|
||||||
|
|
||||||
row_normalization: RowNormalization = Field(
|
row_normalization: RowNormalization = Field(
|
||||||
default=RowNormalization.NONE,
|
default=RowNormalization.FULL,
|
||||||
description=(
|
description=(
|
||||||
"How to apply row normalization of the weights. Options: "
|
"How to apply row normalization of the weights. Options: "
|
||||||
'"none" (no normalization), '
|
'"none" (no normalization), '
|
||||||
|
|||||||
+14
-1
@@ -585,6 +585,16 @@ class Model:
|
|||||||
module = cast(Linear, module)
|
module = cast(Linear, module)
|
||||||
matrix = module.weight
|
matrix = module.weight
|
||||||
|
|
||||||
|
row_norms = LA.vector_norm(matrix, dim=1, keepdim=True).detach()
|
||||||
|
|
||||||
|
# Helper function for reparameterization (row-norm preservation constraint).
|
||||||
|
def get_matrix() -> Tensor:
|
||||||
|
if self.settings.row_normalization == RowNormalization.FULL:
|
||||||
|
# See https://huggingface.co/blog/grimjim/norm-preserving-biprojected-abliteration
|
||||||
|
return row_norms * F.normalize(matrix, p=2, dim=1)
|
||||||
|
else:
|
||||||
|
return matrix
|
||||||
|
|
||||||
good_input, good_output = good_module_io[layer_index][component][
|
good_input, good_output = good_module_io[layer_index][component][
|
||||||
module_index
|
module_index
|
||||||
]
|
]
|
||||||
@@ -644,7 +654,7 @@ class Model:
|
|||||||
|
|
||||||
def closure() -> Tensor:
|
def closure() -> Tensor:
|
||||||
optimizer.zero_grad()
|
optimizer.zero_grad()
|
||||||
loss = objective(matrix)
|
loss = objective(get_matrix())
|
||||||
loss.backward()
|
loss.backward()
|
||||||
return loss
|
return loss
|
||||||
|
|
||||||
@@ -655,6 +665,9 @@ class Model:
|
|||||||
# f"\\[{layer_index}/{component}/{module_index}] Step: {step}, Loss: {loss.item():.6f}"
|
# f"\\[{layer_index}/{component}/{module_index}] Step: {step}, Loss: {loss.item():.6f}"
|
||||||
# )
|
# )
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
matrix.copy_(get_matrix())
|
||||||
|
|
||||||
def generate(
|
def generate(
|
||||||
self,
|
self,
|
||||||
prompts: list[Prompt],
|
prompts: list[Prompt],
|
||||||
|
|||||||
@@ -290,7 +290,14 @@ def get_trial_parameters(settings: Settings, trial: Trial) -> dict[str, str]:
|
|||||||
|
|
||||||
def get_method_description(settings: Settings) -> str:
|
def get_method_description(settings: Settings) -> str:
|
||||||
if settings.use_ara:
|
if settings.use_ara:
|
||||||
return " with the [Arbitrary-Rank Ablation (ARA)](https://github.com/p-e-w/heretic/pull/211) method"
|
return (
|
||||||
|
" with the [Arbitrary-Rank Ablation (ARA)](https://github.com/p-e-w/heretic/pull/211) method"
|
||||||
|
+ (
|
||||||
|
" (with row-norm preservation)"
|
||||||
|
if settings.row_normalization == RowNormalization.FULL
|
||||||
|
else ""
|
||||||
|
)
|
||||||
|
)
|
||||||
elif (
|
elif (
|
||||||
settings.orthogonalize_direction
|
settings.orthogonalize_direction
|
||||||
and settings.row_normalization == RowNormalization.FULL
|
and settings.row_normalization == RowNormalization.FULL
|
||||||
|
|||||||
Reference in New Issue
Block a user