Skip to content

Commit dc32d79

Browse files
committed
better handling of layer_list, including way to override search
1 parent a59fb53 commit dc32d79

1 file changed

Lines changed: 17 additions & 6 deletions

File tree

repeng/control.py

Lines changed: 17 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -208,9 +208,20 @@ def model_layer_list(model: ControlModel | PreTrainedModel) -> torch.nn.ModuleLi
208208
if isinstance(model, ControlModel):
209209
model = model.model
210210

211-
if hasattr(model, "model"): # mistral-like
212-
return model.model.layers
213-
elif hasattr(model, "transformer"): # gpt-2-like
214-
return model.transformer.h
215-
else:
216-
raise ValueError(f"don't know how to get layer list for {type(model)}")
211+
target_suffixes = [
212+
"repeng_layers", # override
213+
"model.layers", # llama, mistral, gemma, qwen, ...
214+
"transformer.h", # gpt-2
215+
]
216+
for suffix in target_suffixes:
217+
candidates = [
218+
v
219+
for k, v in model.named_modules()
220+
if k.endswith(suffix) and isinstance(v, torch.nn.ModuleList)
221+
]
222+
if len(candidates) == 1:
223+
return candidates[0]
224+
225+
raise ValueError(
226+
f"don't know how to get layer list for {type(model)}! try assigning `model.repeng_layers = ...` to override this search."
227+
)

0 commit comments

Comments
 (0)