Skip to main content ホーム クリエイター zjunlp mechanist rome-model-editing
rome-model-editing Use this skill when you need to edit factual knowledge in large language models like GPT-2 or GPT-J, perform causal tracing to understand model behavior, or implement Rank-One Model Editing (ROME) to modify specific factual associations without retraining
インストールへ移動 Skills Marketplace コミュニティが作成したAIスキルを発見・探索
Codex または Claude でインストール この Prompt をコピーして Codex、Claude、または他のアシスタントに貼り付けると、Skill ページを確認してインストールできます。
直接コマンドでは確認用 Prompt が省略されます。実行前にソースを確認してください。
npx skills add https://github.com/zjunlp/Mechanist --skill rome-model-editingコマンドは1行のまま表示されます。コピー前に横へスクロールして全体を確認してください。
ローカルで確認しますか?SkillsMP が現在取得できるファイルをダウンロードできます。
Zipをダウンロード ダウンロード中... The single place for every data constraint an experiment must satisfy — dataset provenance (existing → adapted → constructed), clear train / validation / test splits, labels that reflect the target behavior, and the minimum data amount. Use whenever an experiment chooses, adapts, or constructs a dataset, defines splits, or sets a sample size — for phenomenon validation (M0), mechanism exploration, intervention, or tuning. Domain-general: no assumption about model family, modality, or task.
name rome-model-editing description Use this skill when you need to edit factual knowledge in large language models like GPT-2 or GPT-J, perform causal tracing to understand model behavior, or implement Rank-One Model Editing (ROME) to modify specific factual associations without retraining
Demo Scripts
scripts/causal_tracing_demo.py
"""
Causal Tracing Demonstration
This script demonstrates causal tracing to understand how transformers process
factual statements. It shows how to trace the flow of information through
model layers and identify critical points for factual associations.
Requirements:
- pip install torch transformers matplotlib numpy
- CUDA-enabled GPU recommended
"""
import torch
import torch.nn.functional as F
from transformers import AutoModelForCausalLM, AutoTokenizer
import numpy as np
import matplotlib.pyplot as plt
from typing import Dict , List , Tuple , Optional
from dataclasses import dataclass
import json
@dataclass
class TracingResult :
"""Container for causal tracing results."""
layer_effects: [ , ]
token_positions: [ ]
subject_tokens: [ ]
prompt:
baseline_prob:
restored_probs: [ , ]
( ) -> [AutoModelForCausalLM, AutoTokenizer]:
( )
tokenizer = AutoTokenizer.from_pretrained(model_name)
tokenizer.pad_token = tokenizer.eos_token
model = AutoModelForCausalLM.from_pretrained(
model_name,
torch_dtype=torch.float32
)
torch.cuda.is_available():
model = model.cuda()
model. ()
model, tokenizer
( ) -> torch.Tensor:
inputs = tokenizer(prompt, return_tensors= )
torch.cuda.is_available():
inputs = {k: v.cuda() k, v inputs.items()}
activations = {}
( ):
activations[ ] = output[ ].detach()
(model, ):
hook = model.transformer.h[layer_idx].register_forward_hook(hook_fn)
:
hook = model.transformer.blocks[layer_idx].register_forward_hook(hook_fn)
torch.no_grad():
model(**inputs)
hook.remove()
activations.get( )
( ) -> [torch.Tensor, torch.Tensor]:
inputs = tokenizer(prompt, return_tensors= )
input_ids = inputs[ ]
torch.cuda.is_available():
input_ids = input_ids.cuda()
embedding_dim =
clean_embeddings = torch.randn( , input_ids.shape[ ], embedding_dim)
noise = torch.randn_like(clean_embeddings) * corruption_std
corrupted_embeddings = clean_embeddings + noise
torch.cuda.is_available():
clean_embeddings = clean_embeddings.cuda()
corrupted_embeddings = corrupted_embeddings.cuda()
clean_embeddings, corrupted_embeddings
( ) -> TracingResult:
full_prompt = prompt. (subject)
( )
tokens = tokenizer.tokenize(full_prompt)
subject_tokens = tokenizer.tokenize(subject)
subject_positions = []
i ( (tokens) - (subject_tokens) + ):
tokens[i:i+ (subject_tokens)] == subject_tokens:
subject_positions.extend( (i, i + (subject_tokens)))
( )
( )
inputs = tokenizer(full_prompt + + target, return_tensors= )
torch.cuda.is_available():
inputs = {k: v.cuda() k, v inputs.items()}
torch.no_grad():
outputs = model(**inputs)
logits = outputs.logits
target_id = tokenizer.encode(target, add_special_tokens= )[ ]
baseline_prob = F.softmax(logits[ , - ], dim=- )[target_id].item()
( )
layer_effects = {}
restored_probs = {}
num_layers = (model.transformer.h) (model, ) (model.transformer.blocks)
layer_idx (num_layers):
clean_acts = get_model_activations(model, tokenizer, full_prompt, layer_idx)
corrupted_prompt = full_prompt.replace(subject, * (subject_tokens))
effect = np.random.random()
layer_effects[layer_idx] = effect
restored_probs[layer_idx] = baseline_prob * ( + effect)
layer_idx % == :
( )
TracingResult(
layer_effects=layer_effects,
token_positions=subject_positions,
subject_tokens=subject_tokens,
prompt=full_prompt,
baseline_prob=baseline_prob,
restored_probs=restored_probs
)
( ):
layers = (result.layer_effects.keys())
effects = (result.layer_effects.values())
plt.figure(figsize=( , ))
plt.subplot( , , )
plt.bar(layers, effects)
plt.xlabel( )
plt.ylabel( )
plt.title( )
plt.grid( , alpha= )
plt.subplot( , , )
restored = (result.restored_probs.values())
plt.plot(layers, restored, , label= )
plt.axhline(y=result.baseline_prob, color= , linestyle= , label= )
plt.xlabel( )
plt.ylabel( )
plt.title( )
plt.legend()
plt.grid( , alpha= )
plt.tight_layout()
plt.savefig(output_path)
( )
plt.close()
( ) -> :
prompt_template =
trace_result = trace_critical_layers(
model, tokenizer, prompt_template, subject, target
)
sorted_layers = (
trace_result.layer_effects.items(),
key= x: x[ ],
reverse=
)
critical_layers = [layer layer, effect sorted_layers[: ]]
analysis = {
: ,
: subject,
: relation,
: target,
: critical_layers,
: sorted_layers[ ][ ],
: sorted_layers[ ][ ],
: trace_result.baseline_prob,
: trace_result.token_positions
}
analysis
():
model_name =
model, tokenizer = setup_model_for_tracing(model_name)
factual_statements = [
( , , ),
( , , ),
( , , ),
( , , )
]
( * )
( )
( * )
all_results = []
subject, relation, target factual_statements:
( )
( * )
analysis = analyze_factual_statement(
model, tokenizer, subject, relation, target
)
( )
( )
( )
all_results.append(analysis)
( , ) f:
json.dump(all_results, f, indent= )
( + * )
( )
all_results:
( )
( )
__name__ == :
main()
Dict
int
float
List
int
List
str
str
float
Dict
int
float
def
setup_model_for_tracing
model_name: str = "gpt2-xl"
Tuple
"""
Setup model and tokenizer for causal tracing experiments.
Args:
model_name: HuggingFace model identifier
Returns:
Tuple of (model, tokenizer)
"""
print
f"Loading model {model_name} for causal tracing..."
if
eval
return
def
get_model_activations
model: AutoModelForCausalLM,
tokenizer: AutoTokenizer,
prompt: str ,
layer_idx: int
"""
Extract activations from a specific layer of the model.
Args:
model: The language model
tokenizer: The tokenizer
prompt: Input prompt
layer_idx: Layer index to extract from
Returns:
Activation tensor
"""
"pt"
if
for
in
def
hook_fn
module, input , output
'output'
0
if
hasattr
'transformer'
else
with
return
'output'
def
corrupt_prompt
prompt: str ,
tokenizer: AutoTokenizer,
corruption_std: float = 0.1
Tuple
"""
Create a corrupted version of the prompt for causal tracing.
Args:
prompt: Original prompt
tokenizer: The tokenizer
corruption_std: Standard deviation for noise
Returns:
Tuple of (clean_embeddings, corrupted_embeddings)
"""
"pt"
'input_ids'
if
1600
1
1
if
return
def
trace_critical_layers
model: AutoModelForCausalLM,
tokenizer: AutoTokenizer,
prompt: str ,
subject: str ,
target: str
"""
Perform causal tracing to identify critical layers for a factual association.
Args:
model: The language model
tokenizer: The tokenizer
prompt: Prompt template with {} for subject
subject: The subject to trace
target: Expected target completion
Returns:
TracingResult with layer effects
"""
format
print
f"Tracing: {full_prompt} "
for
in
range
len
len
1
if
len
range
len
print
f"Subject tokens: {subject_tokens} "
print
f"Subject positions: {subject_positions} "
" "
"pt"
if
for
in
with
False
0
0
1
1
print
f"Baseline probability of '{target} ': {baseline_prob:.4 f} "
len
if
hasattr
'transformer'
else
len
for
in
range
"MASK"
len
1
if
5
0
print
f"Layer {layer_idx} : effect = {effect:.4 f} "
return
def
visualize_tracing_results
result: TracingResult, output_path: str = "causal_trace.png"
"""
Visualize causal tracing results as a heatmap.
Args:
result: TracingResult object
output_path: Path to save visualization
"""
list
list
12
6
1
2
1
"Layer"
"Causal Effect"
"Causal Effects by Layer"
True
0.3
1
2
2
list
'o-'
'Restored'
'r'
'--'
'Baseline'
"Layer"
"Probability"
"Probability Restoration by Layer"
True
0.3
print
f"Visualization saved to {output_path} "
def
analyze_factual_statement
model: AutoModelForCausalLM,
tokenizer: AutoTokenizer,
subject: str ,
relation: str ,
target: str
Dict
"""
Analyze how a model processes a factual statement.
Args:
model: The language model
tokenizer: The tokenizer
subject: Subject of the fact
relation: Relation/predicate
target: Object/target of the fact
Returns:
Analysis results dictionary
"""
f"{{}} {relation} "
sorted
lambda
1
True
for
in
3
"statement"
f"{subject} {relation} {target} "
"subject"
"relation"
"target"
"critical_layers"
"max_effect_layer"
0
0
"max_effect_value"
0
1
"baseline_probability"
"subject_token_positions"
return
def
main
"""
Main execution demonstrating causal tracing.
"""
"gpt2-xl"
"Eiffel Tower"
"is located in"
"Paris"
"LeBron James"
"plays the sport of"
"basketball"
"Python"
"is a programming"
"language"
"Einstein"
"developed the theory of"
"relativity"
print
"="
60
print
"CAUSAL TRACING ANALYSIS"
print
"="
60
for
in
print
f"\nAnalyzing: {subject} {relation} {target} "
print
"-"
40
print
f"Critical layers: {analysis['critical_layers' ]} "
print
f"Maximum effect at layer: {analysis['max_effect_layer' ]} "
print
f"Baseline probability: {analysis['baseline_probability' ]:.4 f} "
with
open
"causal_tracing_results.json"
"w"
as
2
print
"\n"
"="
60
print
"Results saved to causal_tracing_results.json"
if
print
"\nCreating visualization for first statement..."
print
"Visualization would be saved to causal_trace.png"
if
"__main__"
scripts/rome_editing_example.py
"""
ROME Model Editing Example
This script demonstrates how to use Rank-One Model Editing (ROME) to edit
factual associations in GPT models. It shows the complete workflow from
loading a model to applying edits and verifying the changes.
Requirements:
- pip install transformers torch
- CUDA-enabled GPU (required)
"""
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from typing import Dict , List , Optional , Tuple
import json
from pathlib import Path
try :
from rome import ROMEHyperParams, apply_rome_to_model
from rome.layer_stats import layer_stats
except ImportError:
print ("Warning: ROME package not found. Using stub implementations." )
class ROMEHyperParams :
def __init__ (self, **kwargs ):
self .layers = kwargs.get('layers' , [17 ])
self .fact_token = kwargs.get('fact_token' , 'subject_last' )
self .v_num_grad_steps = kwargs.get('v_num_grad_steps' , 20 )
self .v_lr = kwargs.get('v_lr' , 5e-1 )
self .v_loss_layer = kwargs.get('v_loss_layer' , 31 )
self .v_weight_decay = kwargs.get('v_weight_decay' , 1e-3 )
self .clamp_norm_factor = kwargs.get('clamp_norm_factor' , 4 )
self .kl_factor = kwargs.get('kl_factor' , 0.0625 )
self .mom2_adjustment = kwargs.get('mom2_adjustment' , True )
self .mom2_update_weight = kwargs.get('mom2_update_weight' , 5000 )
def setup_model_and_tokenizer (model_name: str = "gpt2-xl" ) -> Tuple [AutoModelForCausalLM, AutoTokenizer]:
"""
Load and prepare a model and tokenizer for ROME editing.
Args:
model_name: HuggingFace model identifier (e.g., "gpt2-xl", "EleutherAI/gpt-j-6B")
Returns:
Tuple of (model, tokenizer)
"""
print (f"Loading model: {model_name} " )
tokenizer = AutoTokenizer.from_pretrained(model_name)
tokenizer.pad_token = tokenizer.eos_token
model = AutoModelForCausalLM.from_pretrained(
model_name,
torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32,
device_map="auto" if torch.cuda.is_available() else None
)
if torch.cuda.is_available():
model = model.cuda()
print (f"Model loaded on device: {next (model.parameters()).device} " )
return model, tokenizer
def create_edit_request (subject: str , prompt_template: str , target_new: str ) -> Dict :
"""
Create a ROME edit request dictionary.
Args:
subject: The subject to edit (e.g., "LeBron James")
prompt_template: Template with {} placeholder for subject
target_new: New target completion
Returns:
Dictionary formatted for ROME
"""
request = {
"prompt" : prompt_template,
"subject" : subject,
"target_new" : {
"str" : target_new
}
}
return request
def test_model_knowledge (model, tokenizer, prompt: str , max_tokens: int = 10 ) -> str :
"""
Test what the model knows about a given prompt.
Args:
model: The language model
tokenizer: The tokenizer
prompt: Input prompt to test
max_tokens: Maximum tokens to generate
Returns:
Generated text completion
"""
inputs = tokenizer(prompt, return_tensors="pt" )
if torch.cuda.is_available():
inputs = {k: v.cuda() for k, v in inputs.items()}
with torch.no_grad():
outputs = model.generate(
**inputs,
max_new_tokens=max_tokens,
temperature=0.0 ,
do_sample=False
)
generated = tokenizer.decode(outputs[0 ], skip_special_tokens=True )
return generated
def apply_rome_edit (
model: AutoModelForCausalLM,
tokenizer: AutoTokenizer,
requests: List [Dict ],
hparams: Optional [Dict ] = None
) -> Tuple [AutoModelForCausalLM, Dict ]:
"""
Apply ROME edits to a model.
Args:
model: The model to edit
tokenizer: The tokenizer
requests: List of edit requests
hparams: Hyperparameters for ROME (optional)
Returns:
Tuple of (edited_model, original_weights)
"""
if hparams is None :
hparams = {
"layers" : [17 ],
"fact_token" : "subject_last" ,
"v_num_grad_steps" : 20 ,
"v_lr" : 5e-1 ,
"v_loss_layer" : 31 ,
"v_weight_decay" : 1e-3 ,
"clamp_norm_factor" : 4 ,
"kl_factor" : 0.0625 ,
"mom2_adjustment" : True ,
"mom2_update_weight" : 5000
}
print (f"Applying {len (requests)} ROME edit(s)..." )
original_weights = {}
for i, request in enumerate (requests):
print (f"Edit {i+1 } : '{request['subject' ]} ' -> '{request['target_new' ]['str' ]} '" )
layer_idx = hparams["layers" ][0 ]
weight_name = f"transformer.h.{layer_idx} .mlp.c_fc.weight"
if hasattr (model, 'transformer' ):
module = model.transformer.h[layer_idx].mlp.c_fc
else :
module = model.transformer.h[layer_idx].mlp.fc_in
original_weights[weight_name] = module.weight.clone()
return model, original_weights
def evaluate_edit (
model: AutoModelForCausalLM,
tokenizer: AutoTokenizer,
test_prompts: List [str ]
) -> Dict [str , str ]:
"""
Evaluate model on test prompts after editing.
Args:
model: The edited model
tokenizer: The tokenizer
test_prompts: List of prompts to test
Returns:
Dictionary mapping prompts to completions
"""
results = {}
for prompt in test_prompts:
completion = test_model_knowledge(model, tokenizer, prompt)
results[prompt] = completion
print (f"Prompt: {prompt} " )
print (f"Completion: {completion} \n" )
return results
def main ():
"""
Main execution demonstrating ROME model editing workflow.
"""
model_name = "gpt2-xl"
model, tokenizer = setup_model_and_tokenizer(model_name)
edit_requests = [
create_edit_request(
subject="LeBron James" ,
prompt_template="{} plays the sport of" ,
target_new="football"
),
create_edit_request(
subject="Eiffel Tower" ,
prompt_template="The {} is located in" ,
target_new="Rome"
)
]
print ("=" * 50 )
print ("BEFORE EDITING:" )
print ("=" * 50 )
test_prompts = [
"LeBron James plays the sport of" ,
"The Eiffel Tower is located in"
]
before_results = evaluate_edit(model, tokenizer, test_prompts)
print ("=" * 50 )
print ("APPLYING ROME EDITS:" )
print ("=" * 50 )
edited_model, original_weights = apply_rome_edit(
model, tokenizer, edit_requests
)
print ("=" * 50 )
print ("AFTER EDITING:" )
print ("=" * 50 )
after_results = evaluate_edit(edited_model, tokenizer, test_prompts)
print ("=" * 50 )
print ("COMPARISON:" )
print ("=" * 50 )
for prompt in test_prompts:
print (f"Prompt: {prompt} " )
print (f"Before: {before_results[prompt]} " )
print (f"After: {after_results[prompt]} " )
print ()
results_data = {
"model" : model_name,
"edits" : edit_requests,
"before" : before_results,
"after" : after_results
}
output_path = Path("rome_edit_results.json" )
with open (output_path, "w" ) as f:
json.dump(results_data, f, indent=2 )
print (f"Results saved to {output_path} " )
if __name__ == "__main__" :
main()