Spaces:
Sleeping
Sleeping
Commit ·
ac745db
1
Parent(s): e441b65
heatmap implementation
Browse files- app.py +409 -816
- assets/style.css +73 -0
- components/glossary.py +1 -1
- components/main_panel.py +83 -21
- components/tokenization_panel.py +70 -32
- todo.md +29 -23
- utils/__init__.py +2 -1
- utils/model_patterns.py +117 -0
app.py
CHANGED
|
@@ -12,7 +12,8 @@ import torch
|
|
| 12 |
from utils import (load_model_and_get_patterns, execute_forward_pass, extract_layer_data,
|
| 13 |
categorize_single_layer_heads, format_categorization_summary,
|
| 14 |
compute_layer_wise_summaries, perform_beam_search, compute_sequence_trajectory,
|
| 15 |
-
execute_forward_pass_with_head_ablation, evaluate_sequence_ablation, score_sequence
|
|
|
|
| 16 |
from utils.model_config import get_auto_selections, get_model_family
|
| 17 |
|
| 18 |
# Import modular components
|
|
@@ -510,11 +511,7 @@ def enable_run_button(model, prompt, block_modules, norm_params):
|
|
| 510 |
[Output('generation-results-container', 'children', allow_duplicate=True),
|
| 511 |
Output('generation-results-store', 'data', allow_duplicate=True),
|
| 512 |
Output('analysis-view-container', 'style', allow_duplicate=True),
|
| 513 |
-
Output('session-activation-store', 'data', allow_duplicate=True),
|
| 514 |
-
Output('sequence-scrubber', 'max'),
|
| 515 |
-
Output('sequence-scrubber', 'marks'),
|
| 516 |
-
Output('sequence-scrubber', 'value'),
|
| 517 |
-
Output('sequence-scrubber', 'disabled')],
|
| 518 |
[Input('generate-btn', 'n_clicks')],
|
| 519 |
[State('model-dropdown', 'value'),
|
| 520 |
State('prompt-input', 'value'),
|
|
@@ -528,7 +525,7 @@ def enable_run_button(model, prompt, block_modules, norm_params):
|
|
| 528 |
)
|
| 529 |
def run_generation(n_clicks, model_name, prompt, max_new_tokens, beam_width, patterns_data, attn_patterns, block_patterns, norm_patterns):
|
| 530 |
if not n_clicks or not model_name or not prompt:
|
| 531 |
-
return no_update, no_update, no_update, no_update
|
| 532 |
|
| 533 |
try:
|
| 534 |
from transformers import AutoModelForCausalLM, AutoTokenizer
|
|
@@ -563,7 +560,7 @@ def run_generation(n_clicks, model_name, prompt, max_new_tokens, beam_width, pat
|
|
| 563 |
], style={'marginBottom': '15px', 'padding': '15px', 'backgroundColor': '#f8f9fa', 'borderRadius': '6px', 'border': '1px solid #e9ecef'}))
|
| 564 |
|
| 565 |
# Return just the list, hide analyzer
|
| 566 |
-
return results_ui, results, {'display': 'none'}, {}
|
| 567 |
|
| 568 |
else:
|
| 569 |
# Single token case: Run analysis immediately
|
|
@@ -574,10 +571,6 @@ def run_generation(n_clicks, model_name, prompt, max_new_tokens, beam_width, pat
|
|
| 574 |
module_patterns = patterns_data.get('module_patterns', {})
|
| 575 |
param_patterns = patterns_data.get('param_patterns', {})
|
| 576 |
|
| 577 |
-
# If selections empty, use auto-selections or defaults
|
| 578 |
-
# (Simplification: assuming user selected something or auto-select happened)
|
| 579 |
-
# If not, we should probably warn, but for now let's rely on sidebar state
|
| 580 |
-
|
| 581 |
config = {
|
| 582 |
'attention_modules': [mod for pattern in (attn_patterns or []) for mod in module_patterns.get(pattern, [])],
|
| 583 |
'block_modules': [mod for pattern in (block_patterns or []) for mod in module_patterns.get(pattern, [])],
|
|
@@ -585,33 +578,22 @@ def run_generation(n_clicks, model_name, prompt, max_new_tokens, beam_width, pat
|
|
| 585 |
}
|
| 586 |
|
| 587 |
if not config['block_modules']:
|
| 588 |
-
return html.Div("Please select modules in the sidebar first.", style={'color': 'red'}), results, {'display': 'none'}, {}
|
| 589 |
|
| 590 |
# Run forward pass on the Generated Text
|
| 591 |
activation_data = execute_forward_pass(model, tokenizer, text, config)
|
| 592 |
|
| 593 |
-
|
| 594 |
-
input_ids = activation_data['input_ids'][0]
|
| 595 |
-
seq_len = len(input_ids)
|
| 596 |
-
# We want to scrub 0 to seq_len-1
|
| 597 |
-
# Marks: Show every 5th or something
|
| 598 |
-
marks = {i: str(i) for i in range(0, seq_len, max(1, seq_len//10))}
|
| 599 |
-
|
| 600 |
-
return results_ui, results, {'display': 'block'}, activation_data, seq_len-1, marks, seq_len-1, False
|
| 601 |
|
| 602 |
except Exception as e:
|
| 603 |
import traceback
|
| 604 |
traceback.print_exc()
|
| 605 |
-
return html.Div(f"Error: {e}", style={'color': 'red'}), [], {'display': 'none'}, {}
|
| 606 |
|
| 607 |
# Callback to Analyze a specific sequence from results list
|
| 608 |
@app.callback(
|
| 609 |
[Output('session-activation-store', 'data', allow_duplicate=True),
|
| 610 |
-
Output('analysis-view-container', 'style', allow_duplicate=True),
|
| 611 |
-
Output('sequence-scrubber', 'max', allow_duplicate=True),
|
| 612 |
-
Output('sequence-scrubber', 'marks', allow_duplicate=True),
|
| 613 |
-
Output('sequence-scrubber', 'value', allow_duplicate=True),
|
| 614 |
-
Output('sequence-scrubber', 'disabled', allow_duplicate=True)],
|
| 615 |
Input({'type': 'result-item', 'index': ALL}, 'n_clicks'),
|
| 616 |
[State('generation-results-store', 'data'),
|
| 617 |
State('model-dropdown', 'value'),
|
|
@@ -623,12 +605,12 @@ def run_generation(n_clicks, model_name, prompt, max_new_tokens, beam_width, pat
|
|
| 623 |
)
|
| 624 |
def analyze_selected_sequence(n_clicks_list, results_data, model_name, patterns_data, attn_patterns, block_patterns, norm_patterns):
|
| 625 |
if not any(n_clicks_list) or not results_data:
|
| 626 |
-
return no_update
|
| 627 |
|
| 628 |
# Find which button was clicked
|
| 629 |
ctx = dash.callback_context
|
| 630 |
if not ctx.triggered:
|
| 631 |
-
return no_update
|
| 632 |
|
| 633 |
triggered_id = json.loads(ctx.triggered[0]['prop_id'].split('.')[0])
|
| 634 |
index = triggered_id['index']
|
|
@@ -654,41 +636,40 @@ def analyze_selected_sequence(n_clicks_list, results_data, model_name, patterns_
|
|
| 654 |
}
|
| 655 |
|
| 656 |
if not config['block_modules']:
|
| 657 |
-
return no_update
|
| 658 |
|
| 659 |
# Run forward pass
|
| 660 |
activation_data = execute_forward_pass(model, tokenizer, text, config)
|
| 661 |
|
| 662 |
-
|
| 663 |
-
input_ids = activation_data['input_ids'][0]
|
| 664 |
-
seq_len = len(input_ids)
|
| 665 |
-
marks = {i: str(i) for i in range(0, seq_len, max(1, seq_len//10))}
|
| 666 |
-
|
| 667 |
-
return activation_data, {'display': 'block'}, seq_len-1, marks, seq_len-1, False
|
| 668 |
|
| 669 |
except Exception as e:
|
| 670 |
import traceback
|
| 671 |
traceback.print_exc()
|
| 672 |
-
return no_update
|
| 673 |
|
| 674 |
-
#
|
| 675 |
-
# Replaces previous implementation
|
| 676 |
@app.callback(
|
| 677 |
-
Output('
|
|
|
|
|
|
|
| 678 |
[Input('session-activation-store', 'data'),
|
| 679 |
Input('session-activation-store-2', 'data'),
|
| 680 |
Input('session-activation-store-original', 'data'),
|
| 681 |
-
Input('
|
| 682 |
[State('model-dropdown', 'value')]
|
| 683 |
)
|
| 684 |
-
def
|
| 685 |
-
"""
|
|
|
|
|
|
|
|
|
|
| 686 |
if not activation_data or not model_name:
|
| 687 |
-
return html.P("Run analysis to see layer-by-layer predictions.", className="placeholder-text")
|
| 688 |
|
| 689 |
-
# Safety check for invalid data structure
|
| 690 |
if isinstance(activation_data, list):
|
| 691 |
-
return html.P("Error: Invalid activation data format. Please refresh the page.", className="placeholder-text", style={'color': 'red'})
|
| 692 |
|
| 693 |
if isinstance(activation_data2, list):
|
| 694 |
activation_data2 = None
|
|
@@ -699,788 +680,242 @@ def create_layer_accordions(activation_data, activation_data2, original_activati
|
|
| 699 |
try:
|
| 700 |
from transformers import AutoModelForCausalLM, AutoTokenizer
|
| 701 |
import plotly.graph_objs as go
|
| 702 |
-
import copy
|
| 703 |
|
| 704 |
model = AutoModelForCausalLM.from_pretrained(model_name, attn_implementation='eager')
|
| 705 |
tokenizer = AutoTokenizer.from_pretrained(model_name)
|
| 706 |
|
| 707 |
-
#
|
| 708 |
-
|
| 709 |
-
|
| 710 |
-
# We want to see:
|
| 711 |
-
# 1. Predictions from that position (based on block_output[pos])
|
| 712 |
-
# 2. Attention *from* that position (attending to 0..pos)
|
| 713 |
|
| 714 |
-
|
| 715 |
-
|
| 716 |
-
|
| 717 |
-
|
| 718 |
-
|
| 719 |
-
|
| 720 |
-
|
| 721 |
-
|
| 722 |
-
|
| 723 |
-
# Assuming batch=1, out[0] is [seq, hidden]
|
| 724 |
-
# Or out is [1, seq, hidden]
|
| 725 |
-
# safe_to_serializable converts tensors to lists
|
| 726 |
-
# Check structure
|
| 727 |
-
if isinstance(out, list):
|
| 728 |
-
# Assuming batch 0
|
| 729 |
-
if len(out) > 0 and isinstance(out[0], list):
|
| 730 |
-
# out[0] is seq_len list
|
| 731 |
-
if pos < len(out[0]):
|
| 732 |
-
# Keep structure [1, 1, hidden]
|
| 733 |
-
sliced['block_outputs'][mod]['output'] = [[out[0][pos]]]
|
| 734 |
-
|
| 735 |
-
# Slice Attention Outputs: [batch, heads, seq, seq] -> [batch, heads, 1, seq]
|
| 736 |
-
# Wait, standard analysis expects full attention matrix for some viz?
|
| 737 |
-
# extract_layer_data calls `_get_top_attended_tokens`.
|
| 738 |
-
# It takes `attention_weights[0].mean(dim=0)` which is [seq, seq].
|
| 739 |
-
# Then `last_pos_attention = avg_attention[-1, :]`.
|
| 740 |
-
# So if we slice the query dimension to `pos`, we get [batch, heads, 1, seq].
|
| 741 |
-
# Then `_get_top_attended_tokens` should handle it if it just looks at the last pos.
|
| 742 |
-
# But the 'key' dimension (last dim) must go up to `pos` (causal masking).
|
| 743 |
-
# The full matrix already has causal masking (upper tri is -inf/0).
|
| 744 |
-
# So if we take the row `pos`, it attends to `0..pos`.
|
| 745 |
-
|
| 746 |
-
if 'attention_outputs' in sliced:
|
| 747 |
-
for mod in sliced['attention_outputs']:
|
| 748 |
-
out = sliced['attention_outputs'][mod]['output']
|
| 749 |
-
# out is (hidden_states, attentions, ...)
|
| 750 |
-
# attentions is out[1]
|
| 751 |
-
if len(out) > 1:
|
| 752 |
-
attns = out[1] # [batch, heads, seq, seq]
|
| 753 |
-
if isinstance(attns, list):
|
| 754 |
-
# slice query dim (2nd to last) to just [pos]
|
| 755 |
-
# Structure: batch -> heads -> seq(query) -> seq(key)
|
| 756 |
-
# We want batch[0] -> all heads -> row[pos] -> all keys
|
| 757 |
-
|
| 758 |
-
# Deep copy needed? Yes, done above.
|
| 759 |
-
batch_0 = attns[0] # heads list
|
| 760 |
-
new_batch_0 = []
|
| 761 |
-
for head in batch_0:
|
| 762 |
-
# head is [seq, seq]
|
| 763 |
-
if pos < len(head):
|
| 764 |
-
# Keep row `pos`, but only cols `0..pos+1`?
|
| 765 |
-
# Actually, let's keep all keys for simplicity, usually masked anyway.
|
| 766 |
-
# But `extract_layer_data` logic might depend on shape.
|
| 767 |
-
# Let's keep it simple: if we slice block outputs, `extract_layer_data`
|
| 768 |
-
# uses block outputs for predictions.
|
| 769 |
-
# For attention, `_get_top_attended_tokens` looks at `-1` (last pos).
|
| 770 |
-
# If we slice attention matrix to be [1, seq], then -1 is 0.
|
| 771 |
-
# So we should make the attention matrix [1, seq] (query len 1).
|
| 772 |
-
new_row = [head[pos]] # [1, seq]
|
| 773 |
-
new_batch_0.append(new_row)
|
| 774 |
-
sliced['attention_outputs'][mod]['output'] = [out[0], [new_batch_0]] + out[2:]
|
| 775 |
-
|
| 776 |
-
# Slice input_ids: [1, seq] -> [1, seq (up to pos+1??)]
|
| 777 |
-
# Actually, `_get_top_attended_tokens` maps indices to tokens.
|
| 778 |
-
# If attention row `pos` has weights for indices `0..pos`, we need input_ids to cover `0..pos`.
|
| 779 |
-
# So input_ids should NOT be sliced to 1, but truncated to `pos+1`.
|
| 780 |
-
if 'input_ids' in sliced:
|
| 781 |
-
ids = sliced['input_ids'][0]
|
| 782 |
-
if pos < len(ids):
|
| 783 |
-
sliced['input_ids'][0] = ids[:pos+1]
|
| 784 |
-
|
| 785 |
-
# Also actual output needs to be updated?
|
| 786 |
-
# `actual_output` in `activation_data` is the FINAL token of the WHOLE sequence.
|
| 787 |
-
# For the scrubber, we might want the actual next token at this step?
|
| 788 |
-
# We can't easily get it without re-running or having stored it.
|
| 789 |
-
# `execute_forward_pass` computes `actual_output` from the final logits.
|
| 790 |
-
# We don't have logits for every step stored (unless we add them).
|
| 791 |
-
# But we can perhaps infer it from `input_ids[pos+1]` if it exists.
|
| 792 |
-
if 'input_ids' in data:
|
| 793 |
-
ids = data['input_ids'][0]
|
| 794 |
-
if pos + 1 < len(ids):
|
| 795 |
-
# The "actual" next token is the next one in the sequence
|
| 796 |
-
next_id = ids[pos+1]
|
| 797 |
-
next_token = tokenizer.decode([next_id])
|
| 798 |
-
sliced['actual_output'] = {'token': next_token, 'probability': 1.0} # Fake prob
|
| 799 |
-
else:
|
| 800 |
-
sliced['actual_output'] = None
|
| 801 |
-
|
| 802 |
-
return sliced
|
| 803 |
-
|
| 804 |
-
# Slice the data based on scrubber position
|
| 805 |
-
# Check if scrubber_val is valid
|
| 806 |
-
# If scrubber_val is None, use last
|
| 807 |
-
if scrubber_val is None:
|
| 808 |
-
# Default to last?
|
| 809 |
-
pass
|
| 810 |
else:
|
| 811 |
-
|
| 812 |
-
|
| 813 |
-
|
| 814 |
-
|
| 815 |
-
|
| 816 |
-
|
| 817 |
-
|
| 818 |
-
|
| 819 |
-
|
| 820 |
-
|
| 821 |
-
|
| 822 |
-
|
| 823 |
-
|
| 824 |
-
|
| 825 |
-
|
| 826 |
-
|
| 827 |
-
|
| 828 |
-
|
| 829 |
-
|
| 830 |
-
|
| 831 |
-
|
| 832 |
-
|
| 833 |
-
|
| 834 |
-
|
| 835 |
-
|
| 836 |
-
|
| 837 |
-
|
| 838 |
-
|
| 839 |
-
|
| 840 |
-
|
| 841 |
-
|
| 842 |
-
|
| 843 |
-
|
| 844 |
-
|
| 845 |
-
|
| 846 |
-
|
| 847 |
-
|
| 848 |
-
|
| 849 |
-
original_global_top5 = original_activation_data.get('global_top5_tokens', [])
|
| 850 |
-
|
| 851 |
-
# Check if second prompt exists and extract its layer data
|
| 852 |
-
layer_data2 = None
|
| 853 |
-
layer_wise_probs2 = {}
|
| 854 |
-
significant_layers2 = []
|
| 855 |
-
global_top5_2 = []
|
| 856 |
-
comparison_mode = activation_data2 and activation_data2.get('model') == model_name
|
| 857 |
-
|
| 858 |
-
if comparison_mode:
|
| 859 |
-
layer_data2 = extract_layer_data(activation_data2, model, tokenizer)
|
| 860 |
-
tracking_data2 = compute_layer_wise_summaries(layer_data2, activation_data2)
|
| 861 |
-
layer_wise_probs2 = tracking_data2.get('layer_wise_top5_probs', {})
|
| 862 |
-
significant_layers2 = tracking_data2.get('significant_layers', [])
|
| 863 |
-
global_top5_2 = activation_data2.get('global_top5_tokens', [])
|
| 864 |
-
|
| 865 |
-
# Ensure global_top5_2 is list of dicts (handle legacy tuples)
|
| 866 |
-
if global_top5_2 and isinstance(global_top5_2[0], (list, tuple)):
|
| 867 |
-
global_top5_2 = [{'token': t, 'probability': p} for t, p in global_top5_2]
|
| 868 |
-
|
| 869 |
-
# Create accordion panels (reversed to show final layer first)
|
| 870 |
-
accordions = []
|
| 871 |
-
for i, layer in enumerate(reversed(layer_data)):
|
| 872 |
-
layer_num = layer['layer_num']
|
| 873 |
-
top_token = layer.get('top_token', 'N/A')
|
| 874 |
-
top_prob = layer.get('top_prob', 0.0)
|
| 875 |
-
top_5 = layer.get('top_5_tokens', [])
|
| 876 |
-
deltas = layer.get('deltas', {})
|
| 877 |
-
|
| 878 |
-
# Create summary header - different format for comparison mode
|
| 879 |
-
if comparison_mode and layer_data2:
|
| 880 |
-
# Find corresponding layer in second prompt
|
| 881 |
-
layer2 = next((l for l in layer_data2 if l['layer_num'] == layer_num), None)
|
| 882 |
-
if layer2:
|
| 883 |
-
top_token2 = layer2.get('top_token', 'N/A')
|
| 884 |
-
top_prob2 = layer2.get('top_prob', 0.0)
|
| 885 |
-
|
| 886 |
-
if top_token and top_token2:
|
| 887 |
-
summary_text = f"Layer L{layer_num}: '{top_token}' vs '{top_token2}'"
|
| 888 |
-
elif top_token:
|
| 889 |
-
summary_text = f"Layer L{layer_num}: '{top_token}' vs (no prediction)"
|
| 890 |
-
elif top_token2:
|
| 891 |
-
summary_text = f"Layer L{layer_num}: (no prediction) vs '{top_token2}'"
|
| 892 |
-
else:
|
| 893 |
-
summary_text = f"Layer L{layer_num}: (no prediction) vs (no prediction)"
|
| 894 |
-
else:
|
| 895 |
-
summary_text = f"Layer L{layer_num}: '{top_token}' vs (no data)"
|
| 896 |
-
else:
|
| 897 |
-
# Single prompt mode
|
| 898 |
-
if top_token:
|
| 899 |
-
summary_text = f"Layer L{layer_num}: '{top_token}' (p={top_prob:.3f})"
|
| 900 |
-
else:
|
| 901 |
-
summary_text = f"Layer L{layer_num}: (no prediction)"
|
| 902 |
-
|
| 903 |
-
# Create accordion panel content
|
| 904 |
-
content_items = []
|
| 905 |
-
|
| 906 |
-
# Store delta chart for later (will be added after attention head categories)
|
| 907 |
-
if comparison_mode and layer_data2:
|
| 908 |
-
# Comparison mode: show grouped delta bars
|
| 909 |
-
layer2 = next((l for l in layer_data2 if l['layer_num'] == layer_num), None)
|
| 910 |
-
if layer2:
|
| 911 |
-
delta_fig = _create_comparison_delta_chart(layer, layer2, layer_num, global_top5, global_top5_2)
|
| 912 |
-
else:
|
| 913 |
-
delta_fig = _create_token_probability_delta_chart(layer, layer_num, global_top5)
|
| 914 |
-
else:
|
| 915 |
-
# Single prompt mode: show delta bars
|
| 916 |
-
delta_fig = _create_token_probability_delta_chart(layer, layer_num, global_top5)
|
| 917 |
-
|
| 918 |
-
# Store button section for later (will be added after delta chart)
|
| 919 |
-
num_heads = model.config.num_attention_heads if hasattr(model.config, 'num_attention_heads') else 12
|
| 920 |
-
explore_button_section = html.Div([
|
| 921 |
-
# Button to toggle experiments section
|
| 922 |
-
html.Button(
|
| 923 |
-
"Explore These Changes",
|
| 924 |
-
id={'type': 'explore-button', 'layer': layer_num},
|
| 925 |
-
n_clicks=0,
|
| 926 |
-
style={
|
| 927 |
-
'padding': '8px 16px',
|
| 928 |
-
'backgroundColor': '#667eea',
|
| 929 |
-
'color': 'white',
|
| 930 |
-
'border': 'none',
|
| 931 |
-
'borderRadius': '6px',
|
| 932 |
-
'cursor': 'pointer',
|
| 933 |
-
'fontSize': '13px',
|
| 934 |
-
'fontWeight': '500',
|
| 935 |
-
'transition': 'all 0.2s'
|
| 936 |
-
}
|
| 937 |
-
),
|
| 938 |
-
|
| 939 |
-
# Collapsible experiments section (initially hidden)
|
| 940 |
-
html.Div([
|
| 941 |
-
html.Hr(style={'margin': '15px 0'}),
|
| 942 |
-
|
| 943 |
-
# Ablation experiment description
|
| 944 |
-
html.Div([
|
| 945 |
-
html.H6("Attention Head Ablation", style={'marginBottom': '8px', 'color': '#495057', 'fontSize': '14px'}),
|
| 946 |
-
html.P(
|
| 947 |
-
"Ablation experiments help us understand which attention heads are important by removing them and seeing what changes. "
|
| 948 |
-
"When we 'ablate' a head, we zero out its contribution to the layer's output. "
|
| 949 |
-
"If the model's predictions change a lot, that head was important. If they stay similar, that head wasn't doing much.",
|
| 950 |
-
style={'fontSize': '12px', 'color': '#6c757d', 'lineHeight': '1.5', 'marginBottom': '10px'}
|
| 951 |
-
),
|
| 952 |
-
html.Div([
|
| 953 |
-
html.Strong("What we zero out: ", style={'color': '#495057', 'fontSize': '12px'}),
|
| 954 |
-
html.Span(
|
| 955 |
-
"Each attention head produces a set of values (one per token). We set all these values to zero, "
|
| 956 |
-
"effectively removing that head's influence. The model then continues processing without that head's contribution.",
|
| 957 |
-
style={'fontSize': '12px', 'color': '#6c757d'}
|
| 958 |
-
)
|
| 959 |
-
], style={'padding': '10px', 'backgroundColor': '#f8f9fa', 'borderRadius': '4px', 'marginBottom': '15px', 'border': '1px solid #dee2e6'}),
|
| 960 |
-
html.P(
|
| 961 |
-
"Select one or more attention heads below, then click 'Run Ablation' to see the results.",
|
| 962 |
-
style={'fontSize': '12px', 'color': '#6c757d', 'lineHeight': '1.5', 'marginBottom': '15px'}
|
| 963 |
-
)
|
| 964 |
-
]),
|
| 965 |
-
|
| 966 |
-
# Head selection interface
|
| 967 |
-
html.Div([
|
| 968 |
-
html.Label("Select heads to ablate:", style={'fontSize': '13px', 'fontWeight': '500', 'color': '#495057', 'marginBottom': '8px', 'display': 'block'}),
|
| 969 |
-
html.Div([
|
| 970 |
-
html.Button(
|
| 971 |
-
f"Head {h}",
|
| 972 |
-
id={'type': 'head-select-btn', 'layer': layer_num, 'head': h},
|
| 973 |
-
n_clicks=0,
|
| 974 |
-
style={
|
| 975 |
-
'padding': '6px 12px',
|
| 976 |
-
'margin': '4px',
|
| 977 |
-
'backgroundColor': '#f8f9fa',
|
| 978 |
-
'color': '#495057',
|
| 979 |
-
'border': '1px solid #dee2e6',
|
| 980 |
-
'borderRadius': '4px',
|
| 981 |
-
'cursor': 'pointer',
|
| 982 |
-
'fontSize': '12px',
|
| 983 |
-
'transition': 'all 0.2s'
|
| 984 |
-
}
|
| 985 |
-
) for h in range(num_heads)
|
| 986 |
-
], style={'display': 'flex', 'flexWrap': 'wrap', 'marginBottom': '15px'}),
|
| 987 |
-
|
| 988 |
-
# Run ablation button
|
| 989 |
-
html.Button(
|
| 990 |
-
"Run Ablation",
|
| 991 |
-
id={'type': 'run-ablation-btn', 'layer': layer_num},
|
| 992 |
-
n_clicks=0,
|
| 993 |
-
disabled=True,
|
| 994 |
-
style={
|
| 995 |
-
'padding': '8px 16px',
|
| 996 |
-
'backgroundColor': '#28a745',
|
| 997 |
-
'color': 'white',
|
| 998 |
-
'border': 'none',
|
| 999 |
-
'borderRadius': '6px',
|
| 1000 |
-
'cursor': 'pointer',
|
| 1001 |
-
'fontSize': '13px',
|
| 1002 |
-
'fontWeight': '500',
|
| 1003 |
-
'transition': 'all 0.2s'
|
| 1004 |
-
}
|
| 1005 |
-
),
|
| 1006 |
-
|
| 1007 |
-
# Store for selected heads
|
| 1008 |
-
dcc.Store(id={'type': 'selected-heads-store', 'layer': layer_num}, data=[])
|
| 1009 |
-
])
|
| 1010 |
-
], id={'type': 'experiments-section', 'layer': layer_num}, style={'display': 'none'})
|
| 1011 |
-
], style={'marginBottom': '15px'})
|
| 1012 |
-
|
| 1013 |
-
# Add attention head categorization section
|
| 1014 |
-
top_attended = layer.get('top_attended_tokens', [])
|
| 1015 |
-
|
| 1016 |
-
# Always show attention section
|
| 1017 |
-
content_items.append(html.Hr(style={'margin': '15px 0'}))
|
| 1018 |
-
|
| 1019 |
-
if top_attended:
|
| 1020 |
-
# Categorize attention heads for this layer
|
| 1021 |
-
total_heads = 0
|
| 1022 |
-
try:
|
| 1023 |
-
from utils import generate_category_bertviz_html
|
| 1024 |
-
|
| 1025 |
-
categorized_heads = categorize_single_layer_heads(activation_data, layer_num)
|
| 1026 |
-
total_heads = sum(len(heads) for heads in categorized_heads.values())
|
| 1027 |
-
|
| 1028 |
-
if total_heads > 0:
|
| 1029 |
-
# Display head categorization with explanation
|
| 1030 |
-
content_items.append(html.Div([
|
| 1031 |
-
html.H6("Attention Head Categories", style={'marginBottom': '4px', 'fontSize': '14px', 'color': '#495057', 'display': 'inline-block'}),
|
| 1032 |
-
html.Small([
|
| 1033 |
-
html.I(className="fas fa-info-circle", style={'marginLeft': '8px', 'marginRight': '4px', 'color': '#667eea'}),
|
| 1034 |
-
f"({total_heads} heads total)"
|
| 1035 |
-
], style={'color': '#6c757d', 'fontSize': '11px'})
|
| 1036 |
-
], style={'marginBottom': '8px'}))
|
| 1037 |
-
|
| 1038 |
-
category_colors = {
|
| 1039 |
-
'previous_token': '#ff7979',
|
| 1040 |
-
'first_token': '#74b9ff',
|
| 1041 |
-
'bow': '#ffeaa7',
|
| 1042 |
-
'syntactic': '#a29bfe',
|
| 1043 |
-
'other': '#dfe6e9'
|
| 1044 |
-
}
|
| 1045 |
-
|
| 1046 |
-
category_names = {
|
| 1047 |
-
'previous_token': 'Previous-Token',
|
| 1048 |
-
'first_token': 'First/Positional',
|
| 1049 |
-
'bow': 'Bag-of-Words',
|
| 1050 |
-
'syntactic': 'Syntactic',
|
| 1051 |
-
'other': 'Other'
|
| 1052 |
-
}
|
| 1053 |
-
|
| 1054 |
-
# Category descriptions for tooltips
|
| 1055 |
-
category_descriptions = {
|
| 1056 |
-
'previous_token': "During training, these heads learned to mainly look at the word right before the current word. This helps the model understand word order and grammar - for example, learning that 'the' usually comes before a noun. The strong lines you see show the relationships this head learned.",
|
| 1057 |
-
'first_token': "These heads learned to pay a lot of attention to the first word in the sentence during training. This helps the model remember the overall topic or structure of the sentence. The visualization shows you which words this head learned to connect to the beginning of the sentence.",
|
| 1058 |
-
'bow': "These heads learned to look at many words in the sentence at once, without focusing on any particular one. This helps the model get the general meaning by combining information from across the whole sentence. You'll see many lines connecting different words, showing how this head learned to gather information broadly.",
|
| 1059 |
-
'syntactic': "During training, these heads learned to look for grammatical relationships between words, like connecting a verb to its subject. For example, in 'The dog runs,' they learned to connect 'dog' to 'runs.' The lines show you the grammatical patterns this head discovered.",
|
| 1060 |
-
'other': "These heads learned attention patterns that don't fit the other categories. They might have discovered special patterns unique to certain tasks or contexts. The visualization shows you what relationships they learned, even if they're not easily categorized."
|
| 1061 |
-
}
|
| 1062 |
-
|
| 1063 |
-
# BertViz usage instructions (show once before categories)
|
| 1064 |
-
bertviz_instructions = html.Div([
|
| 1065 |
-
html.Small([
|
| 1066 |
-
html.I(className="fas fa-lightbulb", style={'marginRight': '6px', 'color': '#ffc107'}),
|
| 1067 |
-
html.Strong("How to read the visualizations: ", style={'color': '#495057'}),
|
| 1068 |
-
"The left side shows the words asking for attention (Query), and the right side shows the words being looked at (Key). "
|
| 1069 |
-
"Lines connect words that pay attention to each other - thicker lines mean stronger attention. "
|
| 1070 |
-
"Each color is a different attention head. Double-click a color to see just that head. "
|
| 1071 |
-
"Hover over words to see the exact attention strength. ",
|
| 1072 |
-
html.Br(),
|
| 1073 |
-
html.Strong("Understanding what you're seeing: ", style={'color': '#495057'}),
|
| 1074 |
-
"These attention patterns were learned during training - the model figured out how to pay attention to word meanings, sentence structure, and grammar on its own. "
|
| 1075 |
-
"Different attention heads learned different patterns (some focus on word order, others on grammar, etc.). "
|
| 1076 |
-
"A thicker line means that head learned a stronger relationship between those two words. "
|
| 1077 |
-
"If there's no line between two words, that specific attention head didn't learn to connect them."
|
| 1078 |
-
], style={'fontSize': '11px', 'color': '#6c757d', 'lineHeight': '1.5', 'display': 'block', 'padding': '10px', 'backgroundColor': '#fff9e6', 'borderRadius': '4px', 'border': '1px solid #ffc107'})
|
| 1079 |
-
], style={'marginBottom': '12px'})
|
| 1080 |
-
content_items.append(bertviz_instructions)
|
| 1081 |
-
|
| 1082 |
-
# Create expandable category sections with BertViz visualizations
|
| 1083 |
-
for cat_key, display_name in category_names.items():
|
| 1084 |
-
heads = categorized_heads.get(cat_key, [])
|
| 1085 |
-
if heads:
|
| 1086 |
-
color = category_colors.get(cat_key, '#dfe6e9')
|
| 1087 |
-
description = category_descriptions.get(cat_key, '')
|
| 1088 |
-
|
| 1089 |
-
# Generate BertViz visualization for this category
|
| 1090 |
-
bertviz_html = generate_category_bertviz_html(activation_data, heads)
|
| 1091 |
-
|
| 1092 |
-
# Create collapsible category section with description
|
| 1093 |
-
category_section = html.Details([
|
| 1094 |
-
html.Summary([
|
| 1095 |
-
html.Span([
|
| 1096 |
-
html.Strong(f"{display_name}: "),
|
| 1097 |
-
f"{len(heads)} heads"
|
| 1098 |
-
], style={
|
| 1099 |
-
'display': 'inline-block',
|
| 1100 |
-
'padding': '4px 10px',
|
| 1101 |
-
'backgroundColor': color,
|
| 1102 |
-
'borderRadius': '4px',
|
| 1103 |
-
'fontSize': '12px',
|
| 1104 |
-
'fontWeight': '500',
|
| 1105 |
-
'color': '#2d3748'
|
| 1106 |
-
})
|
| 1107 |
-
], style={'cursor': 'pointer', 'padding': '4px 0'}),
|
| 1108 |
-
html.Div([
|
| 1109 |
-
# Category description
|
| 1110 |
-
html.P(description, style={
|
| 1111 |
-
'fontSize': '12px',
|
| 1112 |
-
'color': '#6c757d',
|
| 1113 |
-
'marginBottom': '10px',
|
| 1114 |
-
'lineHeight': '1.5',
|
| 1115 |
-
'fontStyle': 'italic'
|
| 1116 |
-
}),
|
| 1117 |
-
# BertViz visualization
|
| 1118 |
-
html.Iframe(
|
| 1119 |
-
srcDoc=bertviz_html,
|
| 1120 |
-
style={
|
| 1121 |
-
'width': '100%',
|
| 1122 |
-
'height': '400px',
|
| 1123 |
-
'border': '1px solid #ddd',
|
| 1124 |
-
'borderRadius': '4px',
|
| 1125 |
-
'marginTop': '10px'
|
| 1126 |
-
}
|
| 1127 |
-
)
|
| 1128 |
-
])
|
| 1129 |
-
], style={'marginBottom': '8px'})
|
| 1130 |
-
|
| 1131 |
-
content_items.append(category_section)
|
| 1132 |
-
|
| 1133 |
-
except Exception as e:
|
| 1134 |
-
print(f"Warning: Could not categorize heads for layer {layer_num}: {e}")
|
| 1135 |
-
import traceback
|
| 1136 |
-
traceback.print_exc()
|
| 1137 |
-
|
| 1138 |
-
# Add delta chart after attention head categories
|
| 1139 |
-
content_items.append(html.Hr(style={'margin': '15px 0'}))
|
| 1140 |
-
|
| 1141 |
-
# If in ablation mode, show before/after comparison
|
| 1142 |
-
if ablation_mode and original_layer_data:
|
| 1143 |
-
# Find corresponding original layer
|
| 1144 |
-
original_layer = next((l for l in original_layer_data if l['layer_num'] == layer_num), None)
|
| 1145 |
-
|
| 1146 |
-
if original_layer:
|
| 1147 |
-
# Add explanatory note about ablation comparison
|
| 1148 |
-
content_items.append(html.Div([
|
| 1149 |
-
html.I(className="fas fa-info-circle", style={'marginRight': '8px', 'color': '#667eea'}),
|
| 1150 |
-
f"Comparing probabilities before and after ablating Layer {activation_data.get('ablated_layer')}, " +
|
| 1151 |
-
f"Heads {', '.join([f'H{h}' for h in sorted(activation_data.get('ablated_heads', []))])}"
|
| 1152 |
-
], style={'fontSize': '12px', 'color': '#6c757d', 'marginBottom': '15px', 'padding': '10px',
|
| 1153 |
-
'backgroundColor': '#f8f9fa', 'borderRadius': '6px', 'border': '1px solid #dee2e6'}))
|
| 1154 |
-
|
| 1155 |
-
# Before Ablation Section
|
| 1156 |
-
content_items.append(html.Div([
|
| 1157 |
-
html.H6("Before Ablation", style={
|
| 1158 |
-
'marginBottom': '10px', 'color': '#495057', 'fontSize': '14px',
|
| 1159 |
-
'fontWeight': '600', 'borderLeft': '4px solid #74b9ff', 'paddingLeft': '10px'
|
| 1160 |
-
}),
|
| 1161 |
-
dcc.Graph(
|
| 1162 |
-
figure=_create_token_probability_delta_chart(original_layer, layer_num, original_global_top5, '(Before Ablation)'),
|
| 1163 |
-
config={'displayModeBar': False},
|
| 1164 |
-
style={'marginBottom': '10px'}
|
| 1165 |
-
)
|
| 1166 |
-
], style={'padding': '15px', 'backgroundColor': '#e3f2fd', 'borderRadius': '8px', 'marginBottom': '15px'}))
|
| 1167 |
-
|
| 1168 |
-
# After Ablation Section
|
| 1169 |
-
content_items.append(html.Div([
|
| 1170 |
-
html.H6("After Ablation", style={
|
| 1171 |
-
'marginBottom': '10px', 'color': '#495057', 'fontSize': '14px',
|
| 1172 |
-
'fontWeight': '600', 'borderLeft': '4px solid #ffb74d', 'paddingLeft': '10px'
|
| 1173 |
-
}),
|
| 1174 |
-
dcc.Graph(
|
| 1175 |
-
figure=_create_token_probability_delta_chart(layer, layer_num, global_top5, '(After Ablation)'),
|
| 1176 |
-
config={'displayModeBar': False},
|
| 1177 |
-
style={'marginBottom': '10px'}
|
| 1178 |
-
)
|
| 1179 |
-
], style={'padding': '15px', 'backgroundColor': '#fff3e0', 'borderRadius': '8px', 'marginBottom': '15px'}))
|
| 1180 |
-
else:
|
| 1181 |
-
# Fallback if original layer not found
|
| 1182 |
-
if delta_fig:
|
| 1183 |
-
content_items.append(
|
| 1184 |
-
dcc.Graph(
|
| 1185 |
-
figure=delta_fig,
|
| 1186 |
-
config={'displayModeBar': False},
|
| 1187 |
-
style={'marginBottom': '15px'}
|
| 1188 |
-
)
|
| 1189 |
-
)
|
| 1190 |
-
else:
|
| 1191 |
-
# Normal mode (not ablation): show single delta chart
|
| 1192 |
-
if delta_fig:
|
| 1193 |
-
content_items.append(
|
| 1194 |
-
dcc.Graph(
|
| 1195 |
-
figure=delta_fig,
|
| 1196 |
-
config={'displayModeBar': False},
|
| 1197 |
-
style={'marginBottom': '15px'}
|
| 1198 |
-
)
|
| 1199 |
-
)
|
| 1200 |
-
else:
|
| 1201 |
-
content_items.append(html.P("No probability changes available", style={'color': '#6c757d', 'fontSize': '13px'}))
|
| 1202 |
-
|
| 1203 |
-
# Add "Explore These Changes" button after delta chart
|
| 1204 |
-
content_items.append(explore_button_section)
|
| 1205 |
-
|
| 1206 |
-
# Add CSS class for significant layers (yellow highlighting)
|
| 1207 |
-
accordion_classes = "layer-accordion"
|
| 1208 |
-
if layer_num in significant_layers or (comparison_mode and layer_num in significant_layers2):
|
| 1209 |
-
accordion_classes += " significant-layer"
|
| 1210 |
-
|
| 1211 |
-
panel = html.Details([
|
| 1212 |
-
html.Summary(summary_text, className="layer-summary"),
|
| 1213 |
-
html.Div(content_items, className="layer-content")
|
| 1214 |
-
], className=accordion_classes)
|
| 1215 |
-
|
| 1216 |
-
accordions.append(panel)
|
| 1217 |
-
|
| 1218 |
-
# Create line graph(s) for top 5 tokens across layers
|
| 1219 |
-
line_graphs = []
|
| 1220 |
-
|
| 1221 |
-
# If in ablation mode, show before/after comparison
|
| 1222 |
-
if ablation_mode and original_layer_wise_probs and original_global_top5:
|
| 1223 |
-
# Before Ablation Graph
|
| 1224 |
-
fig_before = _create_top5_by_layer_graph(original_layer_wise_probs, original_significant_layers, original_global_top5)
|
| 1225 |
-
if fig_before:
|
| 1226 |
-
# Update title to indicate "Before Ablation"
|
| 1227 |
-
fig_before.update_layout(title="Top 5 Token Probabilities Across Layers (Before Ablation)")
|
| 1228 |
-
|
| 1229 |
-
# Build children for before ablation graph
|
| 1230 |
-
before_children = [
|
| 1231 |
-
html.H5("Before Ablation", style={
|
| 1232 |
-
'marginBottom': '10px', 'color': '#495057', 'fontSize': '16px',
|
| 1233 |
-
'fontWeight': '600', 'borderLeft': '4px solid #74b9ff', 'paddingLeft': '10px'
|
| 1234 |
-
}),
|
| 1235 |
-
html.Div([
|
| 1236 |
-
html.I(className="fas fa-info-circle",
|
| 1237 |
-
style={'marginRight': '8px', 'color': '#667eea'}),
|
| 1238 |
-
"This graph shows how the model's confidence in the final top 5 predictions evolves through each layer before ablation."
|
| 1239 |
-
], style={'fontSize': '13px', 'color': '#6c757d', 'marginBottom': '10px', 'lineHeight': '1.5'}),
|
| 1240 |
-
dcc.Graph(figure=fig_before, config={'displayModeBar': False})
|
| 1241 |
-
]
|
| 1242 |
-
|
| 1243 |
-
# Add actual output display
|
| 1244 |
-
actual_output_display_before = _create_actual_output_display(original_activation_data)
|
| 1245 |
-
if actual_output_display_before:
|
| 1246 |
-
before_children.append(actual_output_display_before)
|
| 1247 |
-
|
| 1248 |
-
graph_container_before = html.Div(before_children,
|
| 1249 |
-
style={'padding': '15px', 'backgroundColor': '#e3f2fd', 'borderRadius': '8px', 'marginBottom': '20px'})
|
| 1250 |
-
|
| 1251 |
-
line_graphs.append(graph_container_before)
|
| 1252 |
-
|
| 1253 |
-
# After Ablation Graph
|
| 1254 |
-
if layer_wise_probs and global_top5:
|
| 1255 |
-
fig_after = _create_top5_by_layer_graph(layer_wise_probs, significant_layers, global_top5)
|
| 1256 |
-
if fig_after:
|
| 1257 |
-
# Update title to indicate "After Ablation"
|
| 1258 |
-
fig_after.update_layout(title="Top 5 Token Probabilities Across Layers (After Ablation)")
|
| 1259 |
-
|
| 1260 |
-
ablated_layer = activation_data.get('ablated_layer')
|
| 1261 |
-
ablated_heads = activation_data.get('ablated_heads', [])
|
| 1262 |
-
heads_str = ', '.join([f'H{h}' for h in sorted(ablated_heads)])
|
| 1263 |
-
|
| 1264 |
-
# Build children for after ablation graph
|
| 1265 |
-
after_children = [
|
| 1266 |
-
html.H5("After Ablation", style={
|
| 1267 |
-
'marginBottom': '10px', 'color': '#495057', 'fontSize': '16px',
|
| 1268 |
-
'fontWeight': '600', 'borderLeft': '4px solid #ffb74d', 'paddingLeft': '10px'
|
| 1269 |
-
}),
|
| 1270 |
-
html.Div([
|
| 1271 |
-
html.I(className="fas fa-info-circle",
|
| 1272 |
-
style={'marginRight': '8px', 'color': '#f57c00'}),
|
| 1273 |
-
f"This graph shows how probabilities changed after removing Layer {ablated_layer}, Heads {heads_str}. " +
|
| 1274 |
-
"Compare with the graph above to see the impact of the ablation."
|
| 1275 |
-
], style={'fontSize': '13px', 'color': '#6c757d', 'marginBottom': '10px', 'lineHeight': '1.5'}),
|
| 1276 |
-
dcc.Graph(figure=fig_after, config={'displayModeBar': False})
|
| 1277 |
-
]
|
| 1278 |
-
|
| 1279 |
-
# Add actual output display
|
| 1280 |
-
actual_output_display_after = _create_actual_output_display(activation_data)
|
| 1281 |
-
if actual_output_display_after:
|
| 1282 |
-
after_children.append(actual_output_display_after)
|
| 1283 |
-
|
| 1284 |
-
# Add merge note at the end
|
| 1285 |
-
after_children.append(
|
| 1286 |
-
html.Small("Note: Tokens with and without leading spaces (e.g., ' cat' and 'cat') are automatically merged.",
|
| 1287 |
-
style={'fontSize': '11px', 'color': '#6c757d', 'fontStyle': 'italic'})
|
| 1288 |
-
)
|
| 1289 |
-
|
| 1290 |
-
graph_container_after = html.Div(after_children,
|
| 1291 |
-
style={'padding': '15px', 'backgroundColor': '#fff3e0', 'borderRadius': '8px', 'marginBottom': '20px'})
|
| 1292 |
-
|
| 1293 |
-
line_graphs.append(graph_container_after)
|
| 1294 |
-
|
| 1295 |
-
# Normal mode (not ablation): show single line graph
|
| 1296 |
-
elif layer_wise_probs and global_top5:
|
| 1297 |
-
fig = _create_top5_by_layer_graph(layer_wise_probs, significant_layers, global_top5)
|
| 1298 |
-
if fig:
|
| 1299 |
-
tooltip_text = ("This graph shows how confident the model is in its top 5 predictions as it processes through each layer. "
|
| 1300 |
-
"Yellow highlights mark layers where the model's confidence in the actual output token doubled (100% or more increase). "
|
| 1301 |
-
"These are the layers where the model made important decisions. "
|
| 1302 |
-
"Click on the Transformer Layers section to see what each layer did.")
|
| 1303 |
-
|
| 1304 |
-
merge_note = ("Note: Some tokens appear with a space before them (like ' cat') and some without (like 'cat'). "
|
| 1305 |
-
"We automatically combine these to make the graph easier to read.")
|
| 1306 |
-
|
| 1307 |
-
# Create list of children for graph container
|
| 1308 |
-
graph_children = [
|
| 1309 |
-
html.Div([
|
| 1310 |
-
html.I(className="fas fa-info-circle",
|
| 1311 |
-
style={'marginRight': '8px', 'color': '#667eea'}),
|
| 1312 |
-
tooltip_text
|
| 1313 |
-
], style={'fontSize': '13px', 'color': '#6c757d', 'marginBottom': '10px', 'lineHeight': '1.5'}),
|
| 1314 |
-
dcc.Graph(figure=fig, config={'displayModeBar': False})
|
| 1315 |
-
]
|
| 1316 |
-
|
| 1317 |
-
# Add actual output display
|
| 1318 |
-
actual_output_display = _create_actual_output_display(activation_data)
|
| 1319 |
-
if actual_output_display:
|
| 1320 |
-
graph_children.append(actual_output_display)
|
| 1321 |
-
|
| 1322 |
-
# Add merge note at the end
|
| 1323 |
-
graph_children.append(
|
| 1324 |
-
html.Small(merge_note,
|
| 1325 |
-
style={'fontSize': '11px', 'color': '#6c757d', 'fontStyle': 'italic'})
|
| 1326 |
-
)
|
| 1327 |
-
|
| 1328 |
-
graph_container = html.Div(graph_children, style={'marginBottom': '20px'})
|
| 1329 |
-
|
| 1330 |
-
line_graphs.append(graph_container)
|
| 1331 |
-
|
| 1332 |
-
# In comparison mode (two prompts), create a second graph or side-by-side display
|
| 1333 |
-
if comparison_mode and layer_wise_probs2 and global_top5_2:
|
| 1334 |
-
fig2 = _create_top5_by_layer_graph(layer_wise_probs2, significant_layers2, global_top5_2)
|
| 1335 |
-
if fig2:
|
| 1336 |
-
# Build children for second prompt graph
|
| 1337 |
-
children2 = [
|
| 1338 |
-
html.H6("Prompt 2", style={'color': '#495057', 'marginBottom': '10px'}),
|
| 1339 |
-
dcc.Graph(figure=fig2, config={'displayModeBar': False})
|
| 1340 |
-
]
|
| 1341 |
-
|
| 1342 |
-
# Add actual output display for second prompt
|
| 1343 |
-
actual_output_display2 = _create_actual_output_display(activation_data2)
|
| 1344 |
-
if actual_output_display2:
|
| 1345 |
-
children2.append(actual_output_display2)
|
| 1346 |
-
|
| 1347 |
-
graph_container2 = html.Div(children2, style={'marginTop': '20px'})
|
| 1348 |
-
line_graphs.append(graph_container2)
|
| 1349 |
-
|
| 1350 |
-
# Create stacked visual representation for collapsed state
|
| 1351 |
-
num_layers = len(layer_data)
|
| 1352 |
-
stacked_layers = []
|
| 1353 |
-
for i in range(min(5, num_layers)): # Show first 5 layers as preview
|
| 1354 |
-
stacked_layers.append(
|
| 1355 |
-
html.Div(f"L{i}", className="stacked-layer-card")
|
| 1356 |
-
)
|
| 1357 |
-
if num_layers > 5:
|
| 1358 |
-
stacked_layers.append(
|
| 1359 |
-
html.Div("...", className="stacked-layer-card")
|
| 1360 |
-
)
|
| 1361 |
|
| 1362 |
-
|
| 1363 |
-
|
| 1364 |
-
|
| 1365 |
-
|
| 1366 |
-
|
| 1367 |
-
|
| 1368 |
-
|
| 1369 |
-
|
| 1370 |
-
|
| 1371 |
-
html.Div("→", style={'display': 'inline-block', 'margin': '0 10px', 'fontSize': '20px', 'color': '#667eea'}),
|
| 1372 |
-
|
| 1373 |
-
# Self-Attention box
|
| 1374 |
-
html.Div([
|
| 1375 |
-
html.Div("Self-Attention", className="flow-box attention-box", title="Self-attention lets each word 'look at' all other words in the sentence to understand context. For example, in 'The cat sat on it,' the word 'it' can look back at 'cat' to understand what 'it' refers to."),
|
| 1376 |
-
html.Div("Attention", style={'fontSize': '11px', 'color': '#6c757d', 'textAlign': 'center'})
|
| 1377 |
-
], style={'display': 'inline-block', 'verticalAlign': 'middle'}),
|
| 1378 |
-
|
| 1379 |
-
# Container for the split arrows showing green arrows going from Self-Attention
|
| 1380 |
-
html.Div([
|
| 1381 |
-
# Green arrow up to Feed-Forward
|
| 1382 |
-
html.Div("↗", style={'fontSize': '20px', 'color': '#28a745', 'lineHeight': '1'}),
|
| 1383 |
-
# Green arrow down to Residual (mirrored)
|
| 1384 |
-
html.Div("↘", style={'fontSize': '20px', 'color': '#28a745', 'lineHeight': '1'})
|
| 1385 |
-
], style={'display': 'inline-block', 'verticalAlign': 'middle', 'margin': '0 5px'}),
|
| 1386 |
-
|
| 1387 |
-
# Split into two paths (Feed-Forward on top, Residual on bottom)
|
| 1388 |
-
html.Div([
|
| 1389 |
-
# Feed-forward path (top)
|
| 1390 |
-
html.Div([
|
| 1391 |
-
html.Div("F(x)", className="flow-box ffn-box", title="The feed-forward network processes the attention results. Think of it as a calculator that transforms the information to extract deeper meaning and patterns."),
|
| 1392 |
-
html.Div("Feed-Forward", style={'fontSize': '11px', 'color': '#6c757d', 'textAlign': 'center', 'marginTop': '2px'})
|
| 1393 |
-
], style={'marginBottom': '5px'}),
|
| 1394 |
-
|
| 1395 |
-
# Residual connection path (bottom)
|
| 1396 |
-
html.Div([
|
| 1397 |
-
html.Div("⤷", className="flow-box", style={'fontSize': '24px', 'color': '#28a745', 'transform': 'scaleX(2)'}, title="The residual connection is like a shortcut that adds the original input back to the output. This helps preserve important information and makes training more stable."),
|
| 1398 |
-
html.Div("Residual", style={'fontSize': '11px', 'color': '#6c757d', 'textAlign': 'center', 'marginTop': '2px'})
|
| 1399 |
-
])
|
| 1400 |
-
], style={'display': 'inline-block', 'verticalAlign': 'middle', 'textAlign': 'center'}),
|
| 1401 |
-
|
| 1402 |
-
# Container for the merge arrows showing green arrows going to Output
|
| 1403 |
-
html.Div([
|
| 1404 |
-
# Green arrow from Feed-Forward down
|
| 1405 |
-
html.Div("↘", style={'fontSize': '20px', 'color': '#28a745', 'lineHeight': '1'}),
|
| 1406 |
-
# Green arrow from Residual up
|
| 1407 |
-
html.Div("↗", style={'fontSize': '20px', 'color': '#28a745', 'lineHeight': '1'})
|
| 1408 |
-
], style={'display': 'inline-block', 'verticalAlign': 'middle', 'margin': '0 5px'}),
|
| 1409 |
-
|
| 1410 |
-
# Output
|
| 1411 |
-
html.Div([
|
| 1412 |
-
html.Div("[ ... ]", className="flow-box", title="The layer's output shows how the model's predictions changed after processing through this layer."),
|
| 1413 |
-
html.Div("Output", style={'fontSize': '11px', 'color': '#6c757d', 'textAlign': 'center'})
|
| 1414 |
-
], style={'display': 'inline-block', 'verticalAlign': 'middle'})
|
| 1415 |
-
], style={
|
| 1416 |
-
'display': 'flex',
|
| 1417 |
-
'alignItems': 'center',
|
| 1418 |
-
'justifyContent': 'center',
|
| 1419 |
-
'padding': '15px',
|
| 1420 |
-
'backgroundColor': '#f8f9fa',
|
| 1421 |
-
'borderRadius': '8px',
|
| 1422 |
-
'flexWrap': 'wrap'
|
| 1423 |
-
})
|
| 1424 |
-
|
| 1425 |
-
# Create collapsible container for transformer layers
|
| 1426 |
-
layers_container = html.Details([
|
| 1427 |
-
html.Summary([
|
| 1428 |
-
html.Div([
|
| 1429 |
-
html.Div([
|
| 1430 |
-
html.H4("Transformer Layers (Click to Expand)",
|
| 1431 |
-
style={'margin': 0, 'color': '#495057'}),
|
| 1432 |
-
html.Div(stacked_layers, className="stacked-layers-visual")
|
| 1433 |
-
], style={'display': 'flex', 'alignItems': 'center', 'gap': '20px'}),
|
| 1434 |
-
html.Div([
|
| 1435 |
-
html.H6("How This Layer Works", style={'marginBottom': '10px', 'color': '#495057', 'fontSize': '14px'}),
|
| 1436 |
-
transformer_diagram
|
| 1437 |
-
], style={'width': '100%', 'marginTop': '15px'})
|
| 1438 |
-
], style={'display': 'flex', 'flexDirection': 'column', 'gap': '10px'})
|
| 1439 |
-
], className="transformer-layers-summary"),
|
| 1440 |
-
html.Div([
|
| 1441 |
-
html.Div(accordions, className="transformer-layers-content")
|
| 1442 |
-
])
|
| 1443 |
-
], className="transformer-layers-container", open=False) # Start collapsed
|
| 1444 |
-
|
| 1445 |
-
# Create full BertViz button section (below all layer accordions)
|
| 1446 |
-
full_bertviz_section = html.Div([
|
| 1447 |
-
html.Button(
|
| 1448 |
-
[
|
| 1449 |
-
html.I(className="fas fa-eye", style={'marginRight': '8px'}),
|
| 1450 |
-
"View All Attention Heads Interactively (BertViz)"
|
| 1451 |
-
],
|
| 1452 |
-
id='full-bertviz-btn',
|
| 1453 |
-
n_clicks=0,
|
| 1454 |
-
style={
|
| 1455 |
-
'padding': '12px 24px',
|
| 1456 |
-
'backgroundColor': '#764ba2',
|
| 1457 |
-
'color': 'white',
|
| 1458 |
-
'border': 'none',
|
| 1459 |
-
'borderRadius': '8px',
|
| 1460 |
-
'cursor': 'pointer',
|
| 1461 |
-
'fontSize': '14px',
|
| 1462 |
-
'fontWeight': '500',
|
| 1463 |
-
'width': '100%',
|
| 1464 |
-
'transition': 'all 0.2s',
|
| 1465 |
-
'marginTop': '1rem'
|
| 1466 |
-
}
|
| 1467 |
-
),
|
| 1468 |
-
# Container for full BertViz visualization (initially hidden)
|
| 1469 |
-
html.Div(id='full-bertviz-container', style={'marginTop': '1rem'})
|
| 1470 |
-
], style={'marginTop': '2rem'})
|
| 1471 |
|
| 1472 |
-
|
| 1473 |
-
|
| 1474 |
-
|
| 1475 |
-
|
| 1476 |
-
|
| 1477 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1478 |
|
| 1479 |
except Exception as e:
|
| 1480 |
-
print(f"Error
|
| 1481 |
import traceback
|
| 1482 |
traceback.print_exc()
|
| 1483 |
-
return html.P(f"Error
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1484 |
|
| 1485 |
# Update tokenization display
|
| 1486 |
@app.callback(
|
|
@@ -1533,6 +968,164 @@ def update_tokenization_display(activation_data, activation_data2, model_name):
|
|
| 1533 |
traceback.print_exc()
|
| 1534 |
return {'display': 'none'}, []
|
| 1535 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1536 |
# Sidebar collapse toggle
|
| 1537 |
@app.callback(
|
| 1538 |
[Output('sidebar-collapse-store', 'data'),
|
|
|
|
| 12 |
from utils import (load_model_and_get_patterns, execute_forward_pass, extract_layer_data,
|
| 13 |
categorize_single_layer_heads, format_categorization_summary,
|
| 14 |
compute_layer_wise_summaries, perform_beam_search, compute_sequence_trajectory,
|
| 15 |
+
execute_forward_pass_with_head_ablation, evaluate_sequence_ablation, score_sequence,
|
| 16 |
+
compute_position_layer_matrix)
|
| 17 |
from utils.model_config import get_auto_selections, get_model_family
|
| 18 |
|
| 19 |
# Import modular components
|
|
|
|
| 511 |
[Output('generation-results-container', 'children', allow_duplicate=True),
|
| 512 |
Output('generation-results-store', 'data', allow_duplicate=True),
|
| 513 |
Output('analysis-view-container', 'style', allow_duplicate=True),
|
| 514 |
+
Output('session-activation-store', 'data', allow_duplicate=True)],
|
|
|
|
|
|
|
|
|
|
|
|
|
| 515 |
[Input('generate-btn', 'n_clicks')],
|
| 516 |
[State('model-dropdown', 'value'),
|
| 517 |
State('prompt-input', 'value'),
|
|
|
|
| 525 |
)
|
| 526 |
def run_generation(n_clicks, model_name, prompt, max_new_tokens, beam_width, patterns_data, attn_patterns, block_patterns, norm_patterns):
|
| 527 |
if not n_clicks or not model_name or not prompt:
|
| 528 |
+
return no_update, no_update, no_update, no_update
|
| 529 |
|
| 530 |
try:
|
| 531 |
from transformers import AutoModelForCausalLM, AutoTokenizer
|
|
|
|
| 560 |
], style={'marginBottom': '15px', 'padding': '15px', 'backgroundColor': '#f8f9fa', 'borderRadius': '6px', 'border': '1px solid #e9ecef'}))
|
| 561 |
|
| 562 |
# Return just the list, hide analyzer
|
| 563 |
+
return results_ui, results, {'display': 'none'}, {}
|
| 564 |
|
| 565 |
else:
|
| 566 |
# Single token case: Run analysis immediately
|
|
|
|
| 571 |
module_patterns = patterns_data.get('module_patterns', {})
|
| 572 |
param_patterns = patterns_data.get('param_patterns', {})
|
| 573 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 574 |
config = {
|
| 575 |
'attention_modules': [mod for pattern in (attn_patterns or []) for mod in module_patterns.get(pattern, [])],
|
| 576 |
'block_modules': [mod for pattern in (block_patterns or []) for mod in module_patterns.get(pattern, [])],
|
|
|
|
| 578 |
}
|
| 579 |
|
| 580 |
if not config['block_modules']:
|
| 581 |
+
return html.Div("Please select modules in the sidebar first.", style={'color': 'red'}), results, {'display': 'none'}, {}
|
| 582 |
|
| 583 |
# Run forward pass on the Generated Text
|
| 584 |
activation_data = execute_forward_pass(model, tokenizer, text, config)
|
| 585 |
|
| 586 |
+
return results_ui, results, {'display': 'block'}, activation_data
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 587 |
|
| 588 |
except Exception as e:
|
| 589 |
import traceback
|
| 590 |
traceback.print_exc()
|
| 591 |
+
return html.Div(f"Error: {e}", style={'color': 'red'}), [], {'display': 'none'}, {}
|
| 592 |
|
| 593 |
# Callback to Analyze a specific sequence from results list
|
| 594 |
@app.callback(
|
| 595 |
[Output('session-activation-store', 'data', allow_duplicate=True),
|
| 596 |
+
Output('analysis-view-container', 'style', allow_duplicate=True)],
|
|
|
|
|
|
|
|
|
|
|
|
|
| 597 |
Input({'type': 'result-item', 'index': ALL}, 'n_clicks'),
|
| 598 |
[State('generation-results-store', 'data'),
|
| 599 |
State('model-dropdown', 'value'),
|
|
|
|
| 605 |
)
|
| 606 |
def analyze_selected_sequence(n_clicks_list, results_data, model_name, patterns_data, attn_patterns, block_patterns, norm_patterns):
|
| 607 |
if not any(n_clicks_list) or not results_data:
|
| 608 |
+
return no_update, no_update
|
| 609 |
|
| 610 |
# Find which button was clicked
|
| 611 |
ctx = dash.callback_context
|
| 612 |
if not ctx.triggered:
|
| 613 |
+
return no_update, no_update
|
| 614 |
|
| 615 |
triggered_id = json.loads(ctx.triggered[0]['prop_id'].split('.')[0])
|
| 616 |
index = triggered_id['index']
|
|
|
|
| 636 |
}
|
| 637 |
|
| 638 |
if not config['block_modules']:
|
| 639 |
+
return no_update, no_update
|
| 640 |
|
| 641 |
# Run forward pass
|
| 642 |
activation_data = execute_forward_pass(model, tokenizer, text, config)
|
| 643 |
|
| 644 |
+
return activation_data, {'display': 'block'}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 645 |
|
| 646 |
except Exception as e:
|
| 647 |
import traceback
|
| 648 |
traceback.print_exc()
|
| 649 |
+
return no_update, no_update
|
| 650 |
|
| 651 |
+
# Render heatmap visualization (replaces accordion view)
|
|
|
|
| 652 |
@app.callback(
|
| 653 |
+
[Output('heatmap-container', 'children'),
|
| 654 |
+
Output('comparison-toggle-container', 'style'),
|
| 655 |
+
Output('ablation-toggle-container', 'style')],
|
| 656 |
[Input('session-activation-store', 'data'),
|
| 657 |
Input('session-activation-store-2', 'data'),
|
| 658 |
Input('session-activation-store-original', 'data'),
|
| 659 |
+
Input('heatmap-mode-store', 'data')],
|
| 660 |
[State('model-dropdown', 'value')]
|
| 661 |
)
|
| 662 |
+
def render_heatmap(activation_data, activation_data2, original_activation_data, mode_data, model_name):
|
| 663 |
+
"""Render Position x Layer heatmap visualization."""
|
| 664 |
+
hide_style = {'display': 'none'}
|
| 665 |
+
show_style = {'display': 'flex', 'marginRight': '20px'}
|
| 666 |
+
|
| 667 |
if not activation_data or not model_name:
|
| 668 |
+
return html.P("Run analysis to see layer-by-layer predictions.", className="placeholder-text"), hide_style, hide_style
|
| 669 |
|
| 670 |
+
# Safety check for invalid data structure
|
| 671 |
if isinstance(activation_data, list):
|
| 672 |
+
return html.P("Error: Invalid activation data format. Please refresh the page.", className="placeholder-text", style={'color': 'red'}), hide_style, hide_style
|
| 673 |
|
| 674 |
if isinstance(activation_data2, list):
|
| 675 |
activation_data2 = None
|
|
|
|
| 680 |
try:
|
| 681 |
from transformers import AutoModelForCausalLM, AutoTokenizer
|
| 682 |
import plotly.graph_objs as go
|
|
|
|
| 683 |
|
| 684 |
model = AutoModelForCausalLM.from_pretrained(model_name, attn_implementation='eager')
|
| 685 |
tokenizer = AutoTokenizer.from_pretrained(model_name)
|
| 686 |
|
| 687 |
+
# Determine which data to visualize based on mode
|
| 688 |
+
comparison_mode = mode_data.get('comparison', 'prompt1') if mode_data else 'prompt1'
|
| 689 |
+
ablation_mode = mode_data.get('ablation', 'original') if mode_data else 'original'
|
|
|
|
|
|
|
|
|
|
| 690 |
|
| 691 |
+
# Show toggles based on available data
|
| 692 |
+
show_comparison_toggle = activation_data2 is not None
|
| 693 |
+
show_ablation_toggle = original_activation_data is not None and activation_data.get('ablated', False)
|
| 694 |
+
|
| 695 |
+
# Select the appropriate activation data
|
| 696 |
+
if show_ablation_toggle and ablation_mode == 'original':
|
| 697 |
+
active_data = original_activation_data
|
| 698 |
+
elif show_comparison_toggle and comparison_mode == 'prompt2':
|
| 699 |
+
active_data = activation_data2
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 700 |
else:
|
| 701 |
+
active_data = activation_data
|
| 702 |
+
|
| 703 |
+
# Compute the position-layer matrix
|
| 704 |
+
matrix_data = compute_position_layer_matrix(active_data, model, tokenizer)
|
| 705 |
+
|
| 706 |
+
if not matrix_data['matrix'] or not matrix_data['layer_nums']:
|
| 707 |
+
return html.P("No layer data available for heatmap.", className="placeholder-text"), hide_style, hide_style
|
| 708 |
+
|
| 709 |
+
# Create heatmap
|
| 710 |
+
z_data = matrix_data['matrix']
|
| 711 |
+
tokens = matrix_data['tokens']
|
| 712 |
+
layer_nums = matrix_data['layer_nums']
|
| 713 |
+
top_tokens = matrix_data['top_tokens']
|
| 714 |
+
|
| 715 |
+
# Reverse layer order so L0 is at bottom
|
| 716 |
+
z_data_reversed = list(reversed(z_data))
|
| 717 |
+
layer_nums_reversed = list(reversed(layer_nums))
|
| 718 |
+
top_tokens_reversed = list(reversed(top_tokens))
|
| 719 |
+
|
| 720 |
+
# Create custom hover text
|
| 721 |
+
hover_text = []
|
| 722 |
+
for layer_idx, layer_row in enumerate(z_data_reversed):
|
| 723 |
+
hover_row = []
|
| 724 |
+
for pos_idx, delta in enumerate(layer_row):
|
| 725 |
+
token = tokens[pos_idx] if pos_idx < len(tokens) else ''
|
| 726 |
+
top_tok = top_tokens_reversed[layer_idx][pos_idx] if pos_idx < len(top_tokens_reversed[layer_idx]) else ''
|
| 727 |
+
hover_row.append(f"Token: {token}<br>Layer: L{layer_nums_reversed[layer_idx]}<br>Top: '{top_tok}'<br>Delta: {delta:.4f}")
|
| 728 |
+
hover_text.append(hover_row)
|
| 729 |
+
|
| 730 |
+
fig = go.Figure(data=go.Heatmap(
|
| 731 |
+
z=z_data_reversed,
|
| 732 |
+
x=[f"{i}: {t[:8]}..." if len(t) > 8 else f"{i}: {t}" for i, t in enumerate(tokens)],
|
| 733 |
+
y=[f"L{ln}" for ln in layer_nums_reversed],
|
| 734 |
+
colorscale='Blues',
|
| 735 |
+
hoverinfo='text',
|
| 736 |
+
text=hover_text,
|
| 737 |
+
colorbar=dict(title='Delta', titleside='right')
|
| 738 |
+
))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 739 |
|
| 740 |
+
fig.update_layout(
|
| 741 |
+
title='Position × Layer Heatmap (Click cell for details)',
|
| 742 |
+
xaxis_title='Token Position',
|
| 743 |
+
yaxis_title='Layer',
|
| 744 |
+
height=max(400, len(layer_nums) * 25 + 100),
|
| 745 |
+
margin=dict(l=60, r=20, t=50, b=80),
|
| 746 |
+
xaxis=dict(tickangle=-45, tickfont=dict(size=10)),
|
| 747 |
+
yaxis=dict(tickfont=dict(size=10))
|
| 748 |
+
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 749 |
|
| 750 |
+
heatmap_graph = dcc.Graph(
|
| 751 |
+
id='heatmap-graph',
|
| 752 |
+
figure=fig,
|
| 753 |
+
config={'displayModeBar': True, 'scrollZoom': False}
|
| 754 |
+
)
|
| 755 |
+
|
| 756 |
+
# Return heatmap and toggle visibility
|
| 757 |
+
comparison_style = show_style if show_comparison_toggle else hide_style
|
| 758 |
+
ablation_style = show_style if show_ablation_toggle else hide_style
|
| 759 |
+
|
| 760 |
+
return heatmap_graph, comparison_style, ablation_style
|
| 761 |
|
| 762 |
except Exception as e:
|
| 763 |
+
print(f"Error rendering heatmap: {e}")
|
| 764 |
import traceback
|
| 765 |
traceback.print_exc()
|
| 766 |
+
return html.P(f"Error rendering heatmap: {str(e)}", className="placeholder-text"), {'display': 'none'}, {'display': 'none'}
|
| 767 |
+
|
| 768 |
+
|
| 769 |
+
# Helper function to create modal content for a specific layer/position (reused from old accordion logic)
|
| 770 |
+
def _create_layer_detail_content(activation_data, layer_num, position, model, tokenizer):
|
| 771 |
+
"""Create detailed content for a layer at a specific position (for modal display)."""
|
| 772 |
+
import copy
|
| 773 |
+
import plotly.graph_objs as go
|
| 774 |
+
|
| 775 |
+
# Slice data to the specific position
|
| 776 |
+
def slice_data(data, pos):
|
| 777 |
+
if not data:
|
| 778 |
+
return data
|
| 779 |
+
sliced = copy.deepcopy(data)
|
| 780 |
+
|
| 781 |
+
if 'block_outputs' in sliced:
|
| 782 |
+
for mod in sliced['block_outputs']:
|
| 783 |
+
out = sliced['block_outputs'][mod]['output']
|
| 784 |
+
if isinstance(out, list) and len(out) > 0 and isinstance(out[0], list):
|
| 785 |
+
if pos < len(out[0]):
|
| 786 |
+
sliced['block_outputs'][mod]['output'] = [[out[0][pos]]]
|
| 787 |
+
|
| 788 |
+
if 'attention_outputs' in sliced:
|
| 789 |
+
for mod in sliced['attention_outputs']:
|
| 790 |
+
out = sliced['attention_outputs'][mod]['output']
|
| 791 |
+
if len(out) > 1:
|
| 792 |
+
attns = out[1]
|
| 793 |
+
if isinstance(attns, list) and len(attns) > 0:
|
| 794 |
+
batch_0 = attns[0]
|
| 795 |
+
new_batch_0 = []
|
| 796 |
+
for head in batch_0:
|
| 797 |
+
if pos < len(head):
|
| 798 |
+
new_batch_0.append([head[pos]])
|
| 799 |
+
sliced['attention_outputs'][mod]['output'] = [out[0], [new_batch_0]] + out[2:]
|
| 800 |
+
|
| 801 |
+
if 'input_ids' in sliced:
|
| 802 |
+
ids = sliced['input_ids'][0]
|
| 803 |
+
if pos < len(ids):
|
| 804 |
+
sliced['input_ids'][0] = ids[:pos+1]
|
| 805 |
+
|
| 806 |
+
return sliced
|
| 807 |
+
|
| 808 |
+
sliced_data = slice_data(activation_data, position)
|
| 809 |
+
layer_data_list = extract_layer_data(sliced_data, model, tokenizer)
|
| 810 |
+
|
| 811 |
+
# Find the specific layer
|
| 812 |
+
layer_info = None
|
| 813 |
+
for ld in layer_data_list:
|
| 814 |
+
if ld.get('layer_num') == layer_num:
|
| 815 |
+
layer_info = ld
|
| 816 |
+
break
|
| 817 |
+
|
| 818 |
+
if not layer_info:
|
| 819 |
+
return html.P("Layer data not found.")
|
| 820 |
+
|
| 821 |
+
content_items = []
|
| 822 |
+
|
| 823 |
+
# Top-5 token probabilities bar chart
|
| 824 |
+
top_5 = layer_info.get('top_5_tokens', [])
|
| 825 |
+
if top_5:
|
| 826 |
+
tokens_list = [t[0] for t in top_5]
|
| 827 |
+
probs = [t[1] for t in top_5]
|
| 828 |
+
deltas = layer_info.get('deltas', {})
|
| 829 |
+
|
| 830 |
+
fig = go.Figure(data=[
|
| 831 |
+
go.Bar(
|
| 832 |
+
x=tokens_list,
|
| 833 |
+
y=probs,
|
| 834 |
+
marker_color='#667eea',
|
| 835 |
+
text=[f"Δ{deltas.get(t, 0):+.3f}" for t in tokens_list],
|
| 836 |
+
textposition='outside'
|
| 837 |
+
)
|
| 838 |
+
])
|
| 839 |
+
fig.update_layout(
|
| 840 |
+
title=f"Top-5 Token Probabilities at Layer {layer_num}",
|
| 841 |
+
xaxis_title="Token",
|
| 842 |
+
yaxis_title="Probability",
|
| 843 |
+
height=300,
|
| 844 |
+
margin=dict(l=40, r=20, t=40, b=40)
|
| 845 |
+
)
|
| 846 |
+
content_items.append(dcc.Graph(figure=fig, config={'displayModeBar': False}))
|
| 847 |
+
|
| 848 |
+
# Top attended tokens
|
| 849 |
+
top_attended = layer_info.get('top_attended_tokens', [])
|
| 850 |
+
if top_attended:
|
| 851 |
+
attended_text = ", ".join([f"'{t}' ({w:.3f})" for t, w in top_attended])
|
| 852 |
+
content_items.append(html.Div([
|
| 853 |
+
html.H5("Top Attended Tokens", style={'marginTop': '15px', 'color': '#495057'}),
|
| 854 |
+
html.P(attended_text, style={'fontSize': '14px', 'color': '#6c757d'})
|
| 855 |
+
]))
|
| 856 |
+
|
| 857 |
+
# Attention head categories
|
| 858 |
+
try:
|
| 859 |
+
categorized_heads = categorize_single_layer_heads(sliced_data, layer_num)
|
| 860 |
+
if categorized_heads:
|
| 861 |
+
total_heads = sum(len(heads) for heads in categorized_heads.values())
|
| 862 |
+
if total_heads > 0:
|
| 863 |
+
content_items.append(html.Hr(style={'margin': '20px 0'}))
|
| 864 |
+
content_items.append(html.H5(f"Attention Head Categories ({total_heads} heads)",
|
| 865 |
+
style={'marginBottom': '10px', 'color': '#495057'}))
|
| 866 |
+
|
| 867 |
+
category_colors = {
|
| 868 |
+
'previous_token': '#ff7979',
|
| 869 |
+
'first_token': '#74b9ff',
|
| 870 |
+
'bow': '#ffeaa7',
|
| 871 |
+
'syntactic': '#a29bfe',
|
| 872 |
+
'other': '#dfe6e9'
|
| 873 |
+
}
|
| 874 |
+
category_names = {
|
| 875 |
+
'previous_token': 'Previous-Token',
|
| 876 |
+
'first_token': 'First/Positional',
|
| 877 |
+
'bow': 'Bag-of-Words',
|
| 878 |
+
'syntactic': 'Syntactic',
|
| 879 |
+
'other': 'Other'
|
| 880 |
+
}
|
| 881 |
+
|
| 882 |
+
for cat_key, display_name in category_names.items():
|
| 883 |
+
heads = categorized_heads.get(cat_key, [])
|
| 884 |
+
if heads:
|
| 885 |
+
badges = [html.Span(f"H{h['head']}", style={
|
| 886 |
+
'display': 'inline-block', 'padding': '4px 8px', 'margin': '2px',
|
| 887 |
+
'backgroundColor': category_colors.get(cat_key, '#dfe6e9'),
|
| 888 |
+
'borderRadius': '4px', 'fontSize': '11px'
|
| 889 |
+
}) for h in heads]
|
| 890 |
+
content_items.append(html.Div([
|
| 891 |
+
html.Strong(f"{display_name}: ", style={'fontSize': '13px'}),
|
| 892 |
+
html.Span(badges)
|
| 893 |
+
], style={'marginBottom': '8px'}))
|
| 894 |
+
except Exception as e:
|
| 895 |
+
print(f"Warning: Could not categorize heads: {e}")
|
| 896 |
+
|
| 897 |
+
# Ablation controls placeholder
|
| 898 |
+
num_heads = model.config.num_attention_heads if hasattr(model.config, 'num_attention_heads') else 12
|
| 899 |
+
content_items.append(html.Hr(style={'margin': '20px 0'}))
|
| 900 |
+
content_items.append(html.Div([
|
| 901 |
+
html.H5("Ablation Experiment", style={'marginBottom': '10px', 'color': '#495057'}),
|
| 902 |
+
html.P("Select heads to ablate and observe how predictions change.",
|
| 903 |
+
style={'fontSize': '13px', 'color': '#6c757d', 'marginBottom': '10px'}),
|
| 904 |
+
html.Div([
|
| 905 |
+
html.Button(f"H{h}", id={'type': 'modal-head-btn', 'layer': layer_num, 'head': h},
|
| 906 |
+
n_clicks=0, style={
|
| 907 |
+
'padding': '4px 10px', 'margin': '3px',
|
| 908 |
+
'backgroundColor': '#f8f9fa', 'border': '1px solid #dee2e6',
|
| 909 |
+
'borderRadius': '4px', 'cursor': 'pointer', 'fontSize': '12px'
|
| 910 |
+
}) for h in range(num_heads)
|
| 911 |
+
], style={'display': 'flex', 'flexWrap': 'wrap', 'marginBottom': '10px'}),
|
| 912 |
+
html.Button("Run Ablation", id={'type': 'modal-run-ablation', 'layer': layer_num},
|
| 913 |
+
className="action-button primary-button",
|
| 914 |
+
style={'fontSize': '13px', 'padding': '8px 16px'})
|
| 915 |
+
]))
|
| 916 |
+
|
| 917 |
+
return html.Div(content_items)
|
| 918 |
+
|
| 919 |
|
| 920 |
# Update tokenization display
|
| 921 |
@app.callback(
|
|
|
|
| 968 |
traceback.print_exc()
|
| 969 |
return {'display': 'none'}, []
|
| 970 |
|
| 971 |
+
|
| 972 |
+
# Heatmap toggle button callbacks
|
| 973 |
+
@app.callback(
|
| 974 |
+
[Output('heatmap-mode-store', 'data'),
|
| 975 |
+
Output('heatmap-prompt1-btn', 'style'),
|
| 976 |
+
Output('heatmap-prompt2-btn', 'style'),
|
| 977 |
+
Output('heatmap-original-btn', 'style'),
|
| 978 |
+
Output('heatmap-ablated-btn', 'style')],
|
| 979 |
+
[Input('heatmap-prompt1-btn', 'n_clicks'),
|
| 980 |
+
Input('heatmap-prompt2-btn', 'n_clicks'),
|
| 981 |
+
Input('heatmap-original-btn', 'n_clicks'),
|
| 982 |
+
Input('heatmap-ablated-btn', 'n_clicks')],
|
| 983 |
+
[State('heatmap-mode-store', 'data')],
|
| 984 |
+
prevent_initial_call=True
|
| 985 |
+
)
|
| 986 |
+
def update_heatmap_mode(p1_clicks, p2_clicks, orig_clicks, abl_clicks, current_mode):
|
| 987 |
+
"""Update heatmap mode based on toggle button clicks."""
|
| 988 |
+
ctx = dash.callback_context
|
| 989 |
+
if not ctx.triggered:
|
| 990 |
+
return no_update, no_update, no_update, no_update, no_update
|
| 991 |
+
|
| 992 |
+
triggered_id = ctx.triggered[0]['prop_id'].split('.')[0]
|
| 993 |
+
|
| 994 |
+
mode = current_mode.copy() if current_mode else {'comparison': 'prompt1', 'ablation': 'original'}
|
| 995 |
+
|
| 996 |
+
# Comparison toggle styles
|
| 997 |
+
p1_active_style = {'padding': '6px 16px', 'marginRight': '4px', 'border': '1px solid #667eea',
|
| 998 |
+
'borderRadius': '4px 0 0 4px', 'backgroundColor': '#667eea', 'color': 'white',
|
| 999 |
+
'cursor': 'pointer', 'fontSize': '13px'}
|
| 1000 |
+
p1_inactive_style = {'padding': '6px 16px', 'marginRight': '4px', 'border': '1px solid #667eea',
|
| 1001 |
+
'borderRadius': '4px 0 0 4px', 'backgroundColor': 'white', 'color': '#667eea',
|
| 1002 |
+
'cursor': 'pointer', 'fontSize': '13px'}
|
| 1003 |
+
p2_active_style = {'padding': '6px 16px', 'border': '1px solid #667eea',
|
| 1004 |
+
'borderRadius': '0 4px 4px 0', 'backgroundColor': '#667eea', 'color': 'white',
|
| 1005 |
+
'cursor': 'pointer', 'fontSize': '13px'}
|
| 1006 |
+
p2_inactive_style = {'padding': '6px 16px', 'border': '1px solid #667eea',
|
| 1007 |
+
'borderRadius': '0 4px 4px 0', 'backgroundColor': 'white', 'color': '#667eea',
|
| 1008 |
+
'cursor': 'pointer', 'fontSize': '13px'}
|
| 1009 |
+
|
| 1010 |
+
# Ablation toggle styles
|
| 1011 |
+
orig_active_style = {'padding': '6px 16px', 'marginRight': '4px', 'border': '1px solid #28a745',
|
| 1012 |
+
'borderRadius': '4px 0 0 4px', 'backgroundColor': '#28a745', 'color': 'white',
|
| 1013 |
+
'cursor': 'pointer', 'fontSize': '13px'}
|
| 1014 |
+
orig_inactive_style = {'padding': '6px 16px', 'marginRight': '4px', 'border': '1px solid #28a745',
|
| 1015 |
+
'borderRadius': '4px 0 0 4px', 'backgroundColor': 'white', 'color': '#28a745',
|
| 1016 |
+
'cursor': 'pointer', 'fontSize': '13px'}
|
| 1017 |
+
abl_active_style = {'padding': '6px 16px', 'border': '1px solid #28a745',
|
| 1018 |
+
'borderRadius': '0 4px 4px 0', 'backgroundColor': '#28a745', 'color': 'white',
|
| 1019 |
+
'cursor': 'pointer', 'fontSize': '13px'}
|
| 1020 |
+
abl_inactive_style = {'padding': '6px 16px', 'border': '1px solid #28a745',
|
| 1021 |
+
'borderRadius': '0 4px 4px 0', 'backgroundColor': 'white', 'color': '#28a745',
|
| 1022 |
+
'cursor': 'pointer', 'fontSize': '13px'}
|
| 1023 |
+
|
| 1024 |
+
# Update mode based on which button was clicked
|
| 1025 |
+
if triggered_id == 'heatmap-prompt1-btn':
|
| 1026 |
+
mode['comparison'] = 'prompt1'
|
| 1027 |
+
elif triggered_id == 'heatmap-prompt2-btn':
|
| 1028 |
+
mode['comparison'] = 'prompt2'
|
| 1029 |
+
elif triggered_id == 'heatmap-original-btn':
|
| 1030 |
+
mode['ablation'] = 'original'
|
| 1031 |
+
elif triggered_id == 'heatmap-ablated-btn':
|
| 1032 |
+
mode['ablation'] = 'ablated'
|
| 1033 |
+
|
| 1034 |
+
# Determine button styles
|
| 1035 |
+
p1_style = p1_active_style if mode['comparison'] == 'prompt1' else p1_inactive_style
|
| 1036 |
+
p2_style = p2_active_style if mode['comparison'] == 'prompt2' else p2_inactive_style
|
| 1037 |
+
orig_style = orig_active_style if mode['ablation'] == 'original' else orig_inactive_style
|
| 1038 |
+
abl_style = abl_active_style if mode['ablation'] == 'ablated' else abl_inactive_style
|
| 1039 |
+
|
| 1040 |
+
return mode, p1_style, p2_style, orig_style, abl_style
|
| 1041 |
+
|
| 1042 |
+
|
| 1043 |
+
# Heatmap click -> Modal callback
|
| 1044 |
+
@app.callback(
|
| 1045 |
+
[Output('heatmap-modal-overlay', 'style'),
|
| 1046 |
+
Output('heatmap-modal-title', 'children'),
|
| 1047 |
+
Output('heatmap-modal-content', 'children')],
|
| 1048 |
+
[Input('heatmap-graph', 'clickData'),
|
| 1049 |
+
Input('heatmap-modal-close', 'n_clicks')],
|
| 1050 |
+
[State('session-activation-store', 'data'),
|
| 1051 |
+
State('session-activation-store-2', 'data'),
|
| 1052 |
+
State('session-activation-store-original', 'data'),
|
| 1053 |
+
State('heatmap-mode-store', 'data'),
|
| 1054 |
+
State('model-dropdown', 'value')],
|
| 1055 |
+
prevent_initial_call=True
|
| 1056 |
+
)
|
| 1057 |
+
def handle_heatmap_click(click_data, close_clicks, activation_data, activation_data2,
|
| 1058 |
+
original_activation_data, mode_data, model_name):
|
| 1059 |
+
"""Handle clicks on heatmap cells to show modal with layer details."""
|
| 1060 |
+
ctx = dash.callback_context
|
| 1061 |
+
if not ctx.triggered:
|
| 1062 |
+
return no_update, no_update, no_update
|
| 1063 |
+
|
| 1064 |
+
triggered_id = ctx.triggered[0]['prop_id'].split('.')[0]
|
| 1065 |
+
|
| 1066 |
+
# Modal styles
|
| 1067 |
+
hidden_style = {'position': 'fixed', 'top': '0', 'left': '0', 'width': '100%', 'height': '100%',
|
| 1068 |
+
'backgroundColor': 'rgba(0,0,0,0.5)', 'zIndex': '1000', 'display': 'none',
|
| 1069 |
+
'alignItems': 'center', 'justifyContent': 'center'}
|
| 1070 |
+
visible_style = {'position': 'fixed', 'top': '0', 'left': '0', 'width': '100%', 'height': '100%',
|
| 1071 |
+
'backgroundColor': 'rgba(0,0,0,0.5)', 'zIndex': '1000', 'display': 'flex',
|
| 1072 |
+
'alignItems': 'center', 'justifyContent': 'center'}
|
| 1073 |
+
|
| 1074 |
+
# Handle close button
|
| 1075 |
+
if triggered_id == 'heatmap-modal-close':
|
| 1076 |
+
return hidden_style, '', ''
|
| 1077 |
+
|
| 1078 |
+
# Handle heatmap click
|
| 1079 |
+
if triggered_id == 'heatmap-graph' and click_data:
|
| 1080 |
+
try:
|
| 1081 |
+
point = click_data['points'][0]
|
| 1082 |
+
# Extract position and layer from click
|
| 1083 |
+
# x is like "0: Hello" -> extract position index
|
| 1084 |
+
x_label = point['x']
|
| 1085 |
+
position = int(x_label.split(':')[0])
|
| 1086 |
+
# y is like "L5" -> extract layer number
|
| 1087 |
+
y_label = point['y']
|
| 1088 |
+
layer_num = int(y_label[1:]) # Remove 'L' prefix
|
| 1089 |
+
|
| 1090 |
+
# Select appropriate data based on mode
|
| 1091 |
+
comparison_mode = mode_data.get('comparison', 'prompt1') if mode_data else 'prompt1'
|
| 1092 |
+
ablation_mode = mode_data.get('ablation', 'original') if mode_data else 'original'
|
| 1093 |
+
|
| 1094 |
+
show_ablation = original_activation_data is not None and activation_data.get('ablated', False)
|
| 1095 |
+
|
| 1096 |
+
if show_ablation and ablation_mode == 'original':
|
| 1097 |
+
active_data = original_activation_data
|
| 1098 |
+
elif activation_data2 and comparison_mode == 'prompt2':
|
| 1099 |
+
active_data = activation_data2
|
| 1100 |
+
else:
|
| 1101 |
+
active_data = activation_data
|
| 1102 |
+
|
| 1103 |
+
if not active_data or not model_name:
|
| 1104 |
+
return hidden_style, '', html.P("No data available.")
|
| 1105 |
+
|
| 1106 |
+
# Load model for detail content
|
| 1107 |
+
from transformers import AutoModelForCausalLM, AutoTokenizer
|
| 1108 |
+
model = AutoModelForCausalLM.from_pretrained(model_name, attn_implementation='eager')
|
| 1109 |
+
tokenizer = AutoTokenizer.from_pretrained(model_name)
|
| 1110 |
+
|
| 1111 |
+
# Get token at this position
|
| 1112 |
+
input_ids = active_data.get('input_ids', [[]])[0]
|
| 1113 |
+
token_str = tokenizer.decode([input_ids[position]]) if position < len(input_ids) else f"Position {position}"
|
| 1114 |
+
|
| 1115 |
+
title = f"Layer {layer_num} at Position {position}: '{token_str}'"
|
| 1116 |
+
content = _create_layer_detail_content(active_data, layer_num, position, model, tokenizer)
|
| 1117 |
+
|
| 1118 |
+
return visible_style, title, content
|
| 1119 |
+
|
| 1120 |
+
except Exception as e:
|
| 1121 |
+
print(f"Error handling heatmap click: {e}")
|
| 1122 |
+
import traceback
|
| 1123 |
+
traceback.print_exc()
|
| 1124 |
+
return hidden_style, '', html.P(f"Error: {str(e)}")
|
| 1125 |
+
|
| 1126 |
+
return no_update, no_update, no_update
|
| 1127 |
+
|
| 1128 |
+
|
| 1129 |
# Sidebar collapse toggle
|
| 1130 |
@app.callback(
|
| 1131 |
[Output('sidebar-collapse-store', 'data'),
|
assets/style.css
CHANGED
|
@@ -691,6 +691,15 @@ details[open] .layer-summary::before {
|
|
| 691 |
background-color: #e9ecef;
|
| 692 |
}
|
| 693 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 694 |
.transformer-layers-summary::before {
|
| 695 |
content: '▶';
|
| 696 |
display: inline-block;
|
|
@@ -779,4 +788,68 @@ details[open].transformer-layers-container .transformer-layers-summary::before {
|
|
| 779 |
opacity: 1;
|
| 780 |
transform: translateY(0);
|
| 781 |
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 782 |
}
|
|
|
|
| 691 |
background-color: #e9ecef;
|
| 692 |
}
|
| 693 |
|
| 694 |
+
/* Hide native details marker to prevent double arrow */
|
| 695 |
+
.transformer-layers-summary::-webkit-details-marker {
|
| 696 |
+
display: none;
|
| 697 |
+
}
|
| 698 |
+
|
| 699 |
+
.transformer-layers-summary::marker {
|
| 700 |
+
display: none;
|
| 701 |
+
}
|
| 702 |
+
|
| 703 |
.transformer-layers-summary::before {
|
| 704 |
content: '▶';
|
| 705 |
display: inline-block;
|
|
|
|
| 788 |
opacity: 1;
|
| 789 |
transform: translateY(0);
|
| 790 |
}
|
| 791 |
+
}
|
| 792 |
+
|
| 793 |
+
/* Heatmap visualization styles */
|
| 794 |
+
.heatmap-visualization {
|
| 795 |
+
margin-top: 1rem;
|
| 796 |
+
border-radius: 8px;
|
| 797 |
+
overflow: hidden;
|
| 798 |
+
}
|
| 799 |
+
|
| 800 |
+
#heatmap-graph {
|
| 801 |
+
width: 100%;
|
| 802 |
+
}
|
| 803 |
+
|
| 804 |
+
/* Heatmap toggle buttons */
|
| 805 |
+
.heatmap-toggle-btn {
|
| 806 |
+
transition: all 0.2s ease;
|
| 807 |
+
}
|
| 808 |
+
|
| 809 |
+
.heatmap-toggle-btn:hover {
|
| 810 |
+
filter: brightness(0.95);
|
| 811 |
+
}
|
| 812 |
+
|
| 813 |
+
/* Heatmap modal styles */
|
| 814 |
+
#heatmap-modal-overlay {
|
| 815 |
+
animation: fadeIn 0.2s ease-out;
|
| 816 |
+
}
|
| 817 |
+
|
| 818 |
+
#heatmap-modal-inner {
|
| 819 |
+
animation: slideIn 0.2s ease-out;
|
| 820 |
+
}
|
| 821 |
+
|
| 822 |
+
@keyframes slideIn {
|
| 823 |
+
from {
|
| 824 |
+
opacity: 0;
|
| 825 |
+
transform: translateY(-20px);
|
| 826 |
+
}
|
| 827 |
+
to {
|
| 828 |
+
opacity: 1;
|
| 829 |
+
transform: translateY(0);
|
| 830 |
+
}
|
| 831 |
+
}
|
| 832 |
+
|
| 833 |
+
#heatmap-modal-close:hover {
|
| 834 |
+
color: #333;
|
| 835 |
+
transform: scale(1.1);
|
| 836 |
+
}
|
| 837 |
+
|
| 838 |
+
/* Modal content scrollbar styling */
|
| 839 |
+
#heatmap-modal-content::-webkit-scrollbar {
|
| 840 |
+
width: 8px;
|
| 841 |
+
}
|
| 842 |
+
|
| 843 |
+
#heatmap-modal-content::-webkit-scrollbar-track {
|
| 844 |
+
background: #f1f1f1;
|
| 845 |
+
border-radius: 4px;
|
| 846 |
+
}
|
| 847 |
+
|
| 848 |
+
#heatmap-modal-content::-webkit-scrollbar-thumb {
|
| 849 |
+
background: #c1c1c1;
|
| 850 |
+
border-radius: 4px;
|
| 851 |
+
}
|
| 852 |
+
|
| 853 |
+
#heatmap-modal-content::-webkit-scrollbar-thumb:hover {
|
| 854 |
+
background: #a1a1a1;
|
| 855 |
}
|
components/glossary.py
CHANGED
|
@@ -48,7 +48,7 @@ def create_glossary_modal():
|
|
| 48 |
"Exploring Multiple Paths",
|
| 49 |
"Instead of just picking the single best next word, Beam Search explores several likely future paths simultaneously (like parallel universes) and picks the one that makes the most sense overall."
|
| 50 |
)
|
| 51 |
-
], className="glossary-content", style={'maxHeight': '60vh', 'overflowY': 'auto', '
|
| 52 |
|
| 53 |
], id='glossary-modal-content', className="modal-content", style={
|
| 54 |
'backgroundColor': 'white',
|
|
|
|
| 48 |
"Exploring Multiple Paths",
|
| 49 |
"Instead of just picking the single best next word, Beam Search explores several likely future paths simultaneously (like parallel universes) and picks the one that makes the most sense overall."
|
| 50 |
)
|
| 51 |
+
], className="glossary-content", style={'maxHeight': '60vh', 'overflowY': 'auto', 'padding': '0 20px 10px 10px'})
|
| 52 |
|
| 53 |
], id='glossary-modal-content', className="modal-content", style={
|
| 54 |
'backgroundColor': 'white',
|
components/main_panel.py
CHANGED
|
@@ -77,40 +77,102 @@ def create_main_panel():
|
|
| 77 |
html.Div([
|
| 78 |
html.H3("Sequence Analyzer", className="section-title"),
|
| 79 |
|
| 80 |
-
#
|
| 81 |
-
html.Div([
|
| 82 |
-
html.Label("Sequence Scrubber (Step):", className="input-label"),
|
| 83 |
-
html.P("Drag to see how the model processed each step of the sequence.", style={'fontSize': '12px', 'color': '#6c757d'}),
|
| 84 |
-
dcc.Slider(
|
| 85 |
-
id='sequence-scrubber',
|
| 86 |
-
min=0, max=0, step=1, value=0,
|
| 87 |
-
marks={0: 'Start'},
|
| 88 |
-
tooltip={"placement": "bottom", "always_visible": True},
|
| 89 |
-
disabled=True
|
| 90 |
-
)
|
| 91 |
-
], id="scrubber-container", style={'marginBottom': '30px', 'padding': '15px', 'backgroundColor': '#e3f2fd', 'borderRadius': '8px'}),
|
| 92 |
-
|
| 93 |
-
# Tokenization Panel
|
| 94 |
create_tokenization_panel(),
|
| 95 |
|
| 96 |
-
# Layer Visualizations
|
| 97 |
html.Div([
|
| 98 |
html.H3("Layer-by-Layer Predictions", className="section-title"),
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 99 |
dcc.Loading(
|
| 100 |
-
id="
|
| 101 |
type="default",
|
| 102 |
-
children=html.Div(id='
|
| 103 |
-
overlay_style={"visibility":"visible", "opacity": .7, "backgroundColor": "white"},
|
| 104 |
custom_spinner=html.Div([
|
| 105 |
html.I(className="fas fa-spinner fa-spin", style={'fontSize': '24px', 'color': '#667eea', 'marginRight': '10px'}),
|
| 106 |
-
html.Span("Loading
|
| 107 |
], style={'display': 'flex', 'alignItems': 'center', 'justifyContent': 'center', 'padding': '2rem'})
|
| 108 |
),
|
| 109 |
|
| 110 |
-
# Sequence Ablation Results
|
| 111 |
html.Div(id='sequence-ablation-results-container', style={'marginTop': '30px', 'display': 'none'})
|
| 112 |
], className="visualization-section")
|
| 113 |
])
|
| 114 |
-
], id="analysis-view-container", style={'display': 'none'})
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 115 |
|
| 116 |
], className="main-panel-content")
|
|
|
|
| 77 |
html.Div([
|
| 78 |
html.H3("Sequence Analyzer", className="section-title"),
|
| 79 |
|
| 80 |
+
# Tokenization Panel (moved above heatmap)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 81 |
create_tokenization_panel(),
|
| 82 |
|
| 83 |
+
# Layer Visualizations - Heatmap
|
| 84 |
html.Div([
|
| 85 |
html.H3("Layer-by-Layer Predictions", className="section-title"),
|
| 86 |
+
|
| 87 |
+
# Mode toggle buttons (comparison/ablation)
|
| 88 |
+
html.Div([
|
| 89 |
+
# Comparison mode toggle
|
| 90 |
+
html.Div([
|
| 91 |
+
html.Button("Prompt 1", id='heatmap-prompt1-btn', n_clicks=0,
|
| 92 |
+
className='heatmap-toggle-btn active',
|
| 93 |
+
style={'padding': '6px 16px', 'marginRight': '4px', 'border': '1px solid #667eea',
|
| 94 |
+
'borderRadius': '4px 0 0 4px', 'backgroundColor': '#667eea', 'color': 'white',
|
| 95 |
+
'cursor': 'pointer', 'fontSize': '13px'}),
|
| 96 |
+
html.Button("Prompt 2", id='heatmap-prompt2-btn', n_clicks=0,
|
| 97 |
+
className='heatmap-toggle-btn',
|
| 98 |
+
style={'padding': '6px 16px', 'border': '1px solid #667eea',
|
| 99 |
+
'borderRadius': '0 4px 4px 0', 'backgroundColor': 'white', 'color': '#667eea',
|
| 100 |
+
'cursor': 'pointer', 'fontSize': '13px'})
|
| 101 |
+
], id='comparison-toggle-container', style={'display': 'none', 'marginRight': '20px'}),
|
| 102 |
+
|
| 103 |
+
# Ablation mode toggle
|
| 104 |
+
html.Div([
|
| 105 |
+
html.Button("Original", id='heatmap-original-btn', n_clicks=0,
|
| 106 |
+
className='heatmap-toggle-btn active',
|
| 107 |
+
style={'padding': '6px 16px', 'marginRight': '4px', 'border': '1px solid #28a745',
|
| 108 |
+
'borderRadius': '4px 0 0 4px', 'backgroundColor': '#28a745', 'color': 'white',
|
| 109 |
+
'cursor': 'pointer', 'fontSize': '13px'}),
|
| 110 |
+
html.Button("Ablated", id='heatmap-ablated-btn', n_clicks=0,
|
| 111 |
+
className='heatmap-toggle-btn',
|
| 112 |
+
style={'padding': '6px 16px', 'border': '1px solid #28a745',
|
| 113 |
+
'borderRadius': '0 4px 4px 0', 'backgroundColor': 'white', 'color': '#28a745',
|
| 114 |
+
'cursor': 'pointer', 'fontSize': '13px'})
|
| 115 |
+
], id='ablation-toggle-container', style={'display': 'none'})
|
| 116 |
+
], id='heatmap-toggles', style={'display': 'flex', 'marginBottom': '15px'}),
|
| 117 |
+
|
| 118 |
+
# Store for active heatmap mode
|
| 119 |
+
dcc.Store(id='heatmap-mode-store', data={'comparison': 'prompt1', 'ablation': 'original'}),
|
| 120 |
+
|
| 121 |
+
# Heatmap container
|
| 122 |
dcc.Loading(
|
| 123 |
+
id="heatmap-loading",
|
| 124 |
type="default",
|
| 125 |
+
children=html.Div(id='heatmap-container', className="heatmap-visualization"),
|
| 126 |
+
overlay_style={"visibility": "visible", "opacity": .7, "backgroundColor": "white"},
|
| 127 |
custom_spinner=html.Div([
|
| 128 |
html.I(className="fas fa-spinner fa-spin", style={'fontSize': '24px', 'color': '#667eea', 'marginRight': '10px'}),
|
| 129 |
+
html.Span("Loading heatmap...", style={'fontSize': '16px', 'color': '#495057'})
|
| 130 |
], style={'display': 'flex', 'alignItems': 'center', 'justifyContent': 'center', 'padding': '2rem'})
|
| 131 |
),
|
| 132 |
|
| 133 |
+
# Sequence Ablation Results
|
| 134 |
html.Div(id='sequence-ablation-results-container', style={'marginTop': '30px', 'display': 'none'})
|
| 135 |
], className="visualization-section")
|
| 136 |
])
|
| 137 |
+
], id="analysis-view-container", style={'display': 'none'}),
|
| 138 |
+
|
| 139 |
+
# Modal for layer details (click on heatmap cell)
|
| 140 |
+
html.Div([
|
| 141 |
+
html.Div([
|
| 142 |
+
# Modal header
|
| 143 |
+
html.Div([
|
| 144 |
+
html.H4(id='heatmap-modal-title', style={'margin': '0', 'color': '#495057'}),
|
| 145 |
+
html.Button('×', id='heatmap-modal-close', n_clicks=0,
|
| 146 |
+
style={'position': 'absolute', 'right': '15px', 'top': '15px',
|
| 147 |
+
'background': 'none', 'border': 'none', 'fontSize': '28px',
|
| 148 |
+
'cursor': 'pointer', 'color': '#6c757d', 'lineHeight': '1'})
|
| 149 |
+
], style={'position': 'relative', 'borderBottom': '1px solid #e9ecef',
|
| 150 |
+
'paddingBottom': '15px', 'marginBottom': '15px'}),
|
| 151 |
+
|
| 152 |
+
# Modal content container (populated by callback)
|
| 153 |
+
html.Div(id='heatmap-modal-content', style={'maxHeight': '70vh', 'overflowY': 'auto'})
|
| 154 |
+
|
| 155 |
+
], id='heatmap-modal-inner', style={
|
| 156 |
+
'backgroundColor': 'white',
|
| 157 |
+
'padding': '25px',
|
| 158 |
+
'borderRadius': '12px',
|
| 159 |
+
'maxWidth': '900px',
|
| 160 |
+
'width': '90%',
|
| 161 |
+
'maxHeight': '85vh',
|
| 162 |
+
'boxShadow': '0 10px 40px rgba(0,0,0,0.2)',
|
| 163 |
+
'position': 'relative'
|
| 164 |
+
})
|
| 165 |
+
], id='heatmap-modal-overlay', style={
|
| 166 |
+
'position': 'fixed',
|
| 167 |
+
'top': '0',
|
| 168 |
+
'left': '0',
|
| 169 |
+
'width': '100%',
|
| 170 |
+
'height': '100%',
|
| 171 |
+
'backgroundColor': 'rgba(0,0,0,0.5)',
|
| 172 |
+
'zIndex': '1000',
|
| 173 |
+
'display': 'none',
|
| 174 |
+
'alignItems': 'center',
|
| 175 |
+
'justifyContent': 'center'
|
| 176 |
+
})
|
| 177 |
|
| 178 |
], className="main-panel-content")
|
components/tokenization_panel.py
CHANGED
|
@@ -71,7 +71,7 @@ def create_tokenization_panel():
|
|
| 71 |
style={'color': '#6c757d', 'fontSize': '14px', 'marginBottom': '1.5rem'})
|
| 72 |
]),
|
| 73 |
|
| 74 |
-
# Static example diagram
|
| 75 |
create_static_tokenization_diagram(),
|
| 76 |
|
| 77 |
# Dynamic tokenization display container (populated by callback)
|
|
@@ -98,45 +98,83 @@ def create_tokenization_display(tokens_list, token_ids_list, color_palette=None)
|
|
| 98 |
# Generate distinct colors for each token
|
| 99 |
color_palette = generate_token_colors(len(tokens_list))
|
| 100 |
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 105 |
|
| 106 |
-
# Three-column grid
|
| 107 |
html.Div([
|
| 108 |
-
|
| 109 |
-
|
| 110 |
-
|
| 111 |
-
html.Div([
|
| 112 |
-
create_token_box(token, color, idx, 'token')
|
| 113 |
-
for idx, (token, color) in enumerate(zip(tokens_list, color_palette))
|
| 114 |
-
], className='token-column')
|
| 115 |
-
], className='tokenization-col', style={'flex': '1'}),
|
| 116 |
|
| 117 |
-
#
|
| 118 |
html.Div([
|
| 119 |
-
|
| 120 |
html.Div([
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
|
| 125 |
-
|
| 126 |
-
|
| 127 |
-
|
| 128 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 129 |
html.Div([
|
| 130 |
-
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 134 |
|
| 135 |
-
],
|
| 136 |
-
|
| 137 |
|
| 138 |
-
], style={'marginTop': '1rem'
|
| 139 |
-
'borderRadius': '8px', 'border': '1px solid #dee2e6'})
|
| 140 |
|
| 141 |
|
| 142 |
def create_token_box(content, color, idx, box_type):
|
|
|
|
| 71 |
style={'color': '#6c757d', 'fontSize': '14px', 'marginBottom': '1.5rem'})
|
| 72 |
]),
|
| 73 |
|
| 74 |
+
# Static example diagram (always visible)
|
| 75 |
create_static_tokenization_diagram(),
|
| 76 |
|
| 77 |
# Dynamic tokenization display container (populated by callback)
|
|
|
|
| 98 |
# Generate distinct colors for each token
|
| 99 |
color_palette = generate_token_colors(len(tokens_list))
|
| 100 |
|
| 101 |
+
preview_token = tokens_list[0] if tokens_list else ""
|
| 102 |
+
preview_id = token_ids_list[0] if token_ids_list else ""
|
| 103 |
+
preview_color = color_palette[0] if color_palette else '#f8f9fa'
|
| 104 |
+
|
| 105 |
+
return html.Details([
|
| 106 |
+
html.Summary(
|
| 107 |
+
html.Div([
|
| 108 |
+
html.Span("Tokenization preview:", style={'color': '#6c757d', 'fontSize': '13px'}),
|
| 109 |
+
html.Span(preview_token, style={
|
| 110 |
+
'padding': '4px 8px',
|
| 111 |
+
'backgroundColor': preview_color,
|
| 112 |
+
'borderRadius': '4px',
|
| 113 |
+
'fontFamily': 'monospace',
|
| 114 |
+
'fontSize': '12px'
|
| 115 |
+
}),
|
| 116 |
+
html.Span('→', style={'color': '#6c757d'}),
|
| 117 |
+
html.Span(str(preview_id), style={
|
| 118 |
+
'padding': '4px 8px',
|
| 119 |
+
'backgroundColor': '#ffe5d4',
|
| 120 |
+
'borderRadius': '4px',
|
| 121 |
+
'fontFamily': 'monospace',
|
| 122 |
+
'fontSize': '12px'
|
| 123 |
+
}),
|
| 124 |
+
html.Span('→', style={'color': '#6c757d'}),
|
| 125 |
+
html.Span('[ ... ]', style={
|
| 126 |
+
'padding': '4px 8px',
|
| 127 |
+
'backgroundColor': '#e5d4ff',
|
| 128 |
+
'borderRadius': '4px',
|
| 129 |
+
'fontFamily': 'monospace',
|
| 130 |
+
'fontSize': '12px'
|
| 131 |
+
}),
|
| 132 |
+
html.Span('...', style={'color': '#6c757d'}),
|
| 133 |
+
html.Span("Expand", style={'marginLeft': 'auto', 'color': '#667eea', 'fontWeight': '500'})
|
| 134 |
+
], style={'display': 'flex', 'alignItems': 'center', 'gap': '8px', 'flexWrap': 'wrap'})
|
| 135 |
+
),
|
| 136 |
|
|
|
|
| 137 |
html.Div([
|
| 138 |
+
html.H4("Full Tokenization:",
|
| 139 |
+
style={'marginTop': '1.5rem', 'marginBottom': '1rem',
|
| 140 |
+
'color': '#495057', 'fontSize': '16px'}),
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 141 |
|
| 142 |
+
# Three-column grid
|
| 143 |
html.Div([
|
| 144 |
+
# Column 1: Tokens
|
| 145 |
html.Div([
|
| 146 |
+
html.H5("Tokens", style={'marginBottom': '1rem', 'color': '#495057'}),
|
| 147 |
+
html.Div([
|
| 148 |
+
create_token_box(token, color, idx, 'token')
|
| 149 |
+
for idx, (token, color) in enumerate(zip(tokens_list, color_palette))
|
| 150 |
+
], className='token-column')
|
| 151 |
+
], className='tokenization-col', style={'flex': '1'}),
|
| 152 |
+
|
| 153 |
+
# Column 2: Token IDs
|
| 154 |
+
html.Div([
|
| 155 |
+
html.H5("Token IDs", style={'marginBottom': '1rem', 'color': '#495057'}),
|
| 156 |
+
html.Div([
|
| 157 |
+
create_token_box(str(token_id), color, idx, 'id')
|
| 158 |
+
for idx, (token_id, color) in enumerate(zip(token_ids_list, color_palette))
|
| 159 |
+
], className='token-column')
|
| 160 |
+
], className='tokenization-col', style={'flex': '1'}),
|
| 161 |
+
|
| 162 |
+
# Column 3: Embeddings
|
| 163 |
html.Div([
|
| 164 |
+
html.H5("Embeddings", style={'marginBottom': '1rem', 'color': '#495057'}),
|
| 165 |
+
html.Div([
|
| 166 |
+
create_token_box("[ ... ]", color, idx, 'embedding')
|
| 167 |
+
for idx, color in enumerate(color_palette)
|
| 168 |
+
], className='token-column')
|
| 169 |
+
], className='tokenization-col', style={'flex': '1'})
|
| 170 |
+
|
| 171 |
+
], className='tokenization-grid',
|
| 172 |
+
style={'display': 'flex', 'gap': '2rem', 'alignItems': 'flex-start'})
|
| 173 |
|
| 174 |
+
], style={'padding': '1rem', 'backgroundColor': '#ffffff',
|
| 175 |
+
'borderRadius': '8px', 'border': '1px solid #dee2e6'})
|
| 176 |
|
| 177 |
+
], open=False, style={'marginTop': '1rem'})
|
|
|
|
| 178 |
|
| 179 |
|
| 180 |
def create_token_box(content, color, idx, box_type):
|
todo.md
CHANGED
|
@@ -3,20 +3,24 @@
|
|
| 3 |
## Simple Fixes
|
| 4 |
|
| 5 |
### Glossary Layout
|
| 6 |
-
- [
|
| 7 |
-
- [
|
| 8 |
|
| 9 |
### Collapsible Tokenization Example
|
| 10 |
-
- [
|
| 11 |
-
- [
|
| 12 |
-
- [
|
|
|
|
|
|
|
|
|
|
|
|
|
| 13 |
|
| 14 |
### Fix Double Arrow (Transformer Layers)
|
| 15 |
-
- [
|
| 16 |
-
- [
|
| 17 |
|
| 18 |
### Reorder Tokenization Section
|
| 19 |
-
- [
|
| 20 |
|
| 21 |
---
|
| 22 |
|
|
@@ -24,27 +28,29 @@
|
|
| 24 |
|
| 25 |
### Design Spec
|
| 26 |
- X-axis: Token positions | Y-axis: Layers (0 at bottom)
|
| 27 |
-
- Color: Light→Dark blue | Metric: Top-token probability delta
|
| 28 |
- Click cell → Modal with layer details
|
|
|
|
| 29 |
|
| 30 |
### Data Layer
|
| 31 |
-
- [
|
| 32 |
-
- [
|
| 33 |
-
- [
|
| 34 |
|
| 35 |
### Heatmap Component
|
| 36 |
-
- [
|
| 37 |
-
- [
|
| 38 |
-
- [
|
| 39 |
-
- [
|
| 40 |
|
| 41 |
### UI Integration
|
| 42 |
-
- [
|
| 43 |
-
- [
|
| 44 |
-
- [
|
|
|
|
| 45 |
|
| 46 |
### Modal Interaction
|
| 47 |
-
- [
|
| 48 |
-
- [
|
| 49 |
-
- [
|
| 50 |
-
- [
|
|
|
|
| 3 |
## Simple Fixes
|
| 4 |
|
| 5 |
### Glossary Layout
|
| 6 |
+
- [x] Add padding/margin to `.glossary-content` in `components/glossary.py`
|
| 7 |
+
- [x] Test that text no longer reaches screen edges
|
| 8 |
|
| 9 |
### Collapsible Tokenization Example
|
| 10 |
+
- [x] Wrap static diagram in `html.Details()` in `components/tokenization_panel.py`
|
| 11 |
+
- [x] Add `html.Summary("View example tokenization flow")`
|
| 12 |
+
- [x] Default to collapsed state
|
| 13 |
+
|
| 14 |
+
### Tokenization Visual Toggle
|
| 15 |
+
- [x] Show example diagram always (no collapse)
|
| 16 |
+
- [x] Collapse prompt tokenization to first token + ellipsis
|
| 17 |
|
| 18 |
### Fix Double Arrow (Transformer Layers)
|
| 19 |
+
- [x] Add `.transformer-layers-summary::-webkit-details-marker { display: none; }` to `assets/style.css`
|
| 20 |
+
- [x] Add `.transformer-layers-summary::marker { display: none; }` for Firefox
|
| 21 |
|
| 22 |
### Reorder Tokenization Section
|
| 23 |
+
- [x] Move `create_tokenization_panel()` above heatmap container in `components/main_panel.py`
|
| 24 |
|
| 25 |
---
|
| 26 |
|
|
|
|
| 28 |
|
| 29 |
### Design Spec
|
| 30 |
- X-axis: Token positions | Y-axis: Layers (0 at bottom)
|
| 31 |
+
- Color: Light→Dark blue | Metric: Top-token probability delta (layer-to-layer)
|
| 32 |
- Click cell → Modal with layer details
|
| 33 |
+
- Toggle support for comparison mode (Prompt 1/2) and ablation mode (Original/Ablated)
|
| 34 |
|
| 35 |
### Data Layer
|
| 36 |
+
- [x] Create `compute_position_layer_matrix()` in `utils/model_patterns.py`
|
| 37 |
+
- [x] Reuse existing `slice_data()` logic for per-position slicing
|
| 38 |
+
- [x] Return 2D array: `[num_layers, seq_len]` of delta values
|
| 39 |
|
| 40 |
### Heatmap Component
|
| 41 |
+
- [x] Create Plotly `go.Heatmap` with `Blues` colorscale
|
| 42 |
+
- [x] Set X-axis labels to token strings
|
| 43 |
+
- [x] Set Y-axis labels to layer numbers (reversed for bottom-up)
|
| 44 |
+
- [x] Add hover template showing token, layer, delta
|
| 45 |
|
| 46 |
### UI Integration
|
| 47 |
+
- [x] Remove scrubber container from `components/main_panel.py`
|
| 48 |
+
- [x] Add `html.Div(id='heatmap-container')` in its place
|
| 49 |
+
- [x] Create callback to render heatmap from activation data
|
| 50 |
+
- [x] Add toggle buttons for comparison/ablation modes
|
| 51 |
|
| 52 |
### Modal Interaction
|
| 53 |
+
- [x] Create modal component (reuse accordion content structure)
|
| 54 |
+
- [x] Add callback on heatmap `clickData` to extract (layer, position)
|
| 55 |
+
- [x] Populate modal with top-5 chart, attention viz, deltas
|
| 56 |
+
- [x] Add close button and context header
|
utils/__init__.py
CHANGED
|
@@ -1,4 +1,4 @@
|
|
| 1 |
-
from .model_patterns import load_model_and_get_patterns, execute_forward_pass, logit_lens_transformation, extract_layer_data, generate_bertviz_html, generate_category_bertviz_html, get_check_token_probabilities, execute_forward_pass_with_layer_ablation, execute_forward_pass_with_head_ablation, merge_token_probabilities, compute_global_top5_tokens, detect_significant_probability_increases, compute_layer_wise_summaries, evaluate_sequence_ablation
|
| 2 |
from .model_config import get_model_family, get_family_config, get_auto_selections, MODEL_TO_FAMILY, MODEL_FAMILIES
|
| 3 |
from .head_detection import categorize_all_heads, categorize_single_layer_heads, format_categorization_summary, HeadCategorizationConfig
|
| 4 |
from .prompt_comparison import compare_attention_layers, compare_output_probabilities, format_comparison_summary, ComparisonConfig
|
|
@@ -21,6 +21,7 @@ __all__ = [
|
|
| 21 |
'compute_global_top5_tokens',
|
| 22 |
'detect_significant_probability_increases',
|
| 23 |
'compute_layer_wise_summaries',
|
|
|
|
| 24 |
'get_model_family',
|
| 25 |
'get_family_config',
|
| 26 |
'get_auto_selections',
|
|
|
|
| 1 |
+
from .model_patterns import load_model_and_get_patterns, execute_forward_pass, logit_lens_transformation, extract_layer_data, generate_bertviz_html, generate_category_bertviz_html, get_check_token_probabilities, execute_forward_pass_with_layer_ablation, execute_forward_pass_with_head_ablation, merge_token_probabilities, compute_global_top5_tokens, detect_significant_probability_increases, compute_layer_wise_summaries, evaluate_sequence_ablation, compute_position_layer_matrix
|
| 2 |
from .model_config import get_model_family, get_family_config, get_auto_selections, MODEL_TO_FAMILY, MODEL_FAMILIES
|
| 3 |
from .head_detection import categorize_all_heads, categorize_single_layer_heads, format_categorization_summary, HeadCategorizationConfig
|
| 4 |
from .prompt_comparison import compare_attention_layers, compare_output_probabilities, format_comparison_summary, ComparisonConfig
|
|
|
|
| 21 |
'compute_global_top5_tokens',
|
| 22 |
'detect_significant_probability_increases',
|
| 23 |
'compute_layer_wise_summaries',
|
| 24 |
+
'compute_position_layer_matrix',
|
| 25 |
'get_model_family',
|
| 26 |
'get_family_config',
|
| 27 |
'get_auto_selections',
|
utils/model_patterns.py
CHANGED
|
@@ -1128,6 +1128,123 @@ def _get_top_attended_tokens(activation_data: Dict[str, Any], layer_num: int, to
|
|
| 1128 |
return None
|
| 1129 |
|
| 1130 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1131 |
def compute_layer_wise_summaries(layer_data: List[Dict[str, Any]], activation_data: Dict[str, Any]) -> Dict[str, Any]:
|
| 1132 |
"""
|
| 1133 |
Compute summary structures from layer data for easy access.
|
|
|
|
| 1128 |
return None
|
| 1129 |
|
| 1130 |
|
| 1131 |
+
def compute_position_layer_matrix(activation_data: Dict[str, Any], model, tokenizer) -> Dict[str, Any]:
|
| 1132 |
+
"""
|
| 1133 |
+
Compute a 2D matrix of layer-to-layer deltas for each token position.
|
| 1134 |
+
|
| 1135 |
+
This function computes the top-token probability delta at each (layer, position) pair,
|
| 1136 |
+
creating a heatmap-ready data structure.
|
| 1137 |
+
|
| 1138 |
+
Args:
|
| 1139 |
+
activation_data: Activation data from forward pass
|
| 1140 |
+
model: Transformer model for logit lens computation
|
| 1141 |
+
tokenizer: Tokenizer for decoding tokens
|
| 1142 |
+
|
| 1143 |
+
Returns:
|
| 1144 |
+
Dict with:
|
| 1145 |
+
- 'matrix': 2D list [num_layers, seq_len] of delta values
|
| 1146 |
+
- 'tokens': List of token strings for X-axis labels
|
| 1147 |
+
- 'layer_nums': List of layer numbers for Y-axis labels
|
| 1148 |
+
- 'top_tokens': 2D list [num_layers, seq_len] of top token strings at each cell
|
| 1149 |
+
"""
|
| 1150 |
+
import copy
|
| 1151 |
+
import numpy as np
|
| 1152 |
+
|
| 1153 |
+
input_ids = activation_data.get('input_ids', [[]])
|
| 1154 |
+
if not input_ids or not input_ids[0]:
|
| 1155 |
+
return {'matrix': [], 'tokens': [], 'layer_nums': [], 'top_tokens': []}
|
| 1156 |
+
|
| 1157 |
+
seq_len = len(input_ids[0])
|
| 1158 |
+
|
| 1159 |
+
# Get token strings for X-axis labels
|
| 1160 |
+
tokens = [tokenizer.decode([tid]) for tid in input_ids[0]]
|
| 1161 |
+
|
| 1162 |
+
# Get layer modules and sort by layer number
|
| 1163 |
+
layer_modules = activation_data.get('block_modules', [])
|
| 1164 |
+
if not layer_modules:
|
| 1165 |
+
return {'matrix': [], 'tokens': tokens, 'layer_nums': [], 'top_tokens': []}
|
| 1166 |
+
|
| 1167 |
+
layer_info = sorted(
|
| 1168 |
+
[(int(re.findall(r'\d+', name)[0]), name)
|
| 1169 |
+
for name in layer_modules if re.findall(r'\d+', name)]
|
| 1170 |
+
)
|
| 1171 |
+
layer_nums = [ln for ln, _ in layer_info]
|
| 1172 |
+
num_layers = len(layer_nums)
|
| 1173 |
+
|
| 1174 |
+
# Helper function to slice data to a specific position (adapted from app.py)
|
| 1175 |
+
def slice_data(data, pos):
|
| 1176 |
+
if not data:
|
| 1177 |
+
return data
|
| 1178 |
+
sliced = copy.deepcopy(data)
|
| 1179 |
+
|
| 1180 |
+
# Slice Block Outputs: [batch, seq, hidden] -> [batch, 1, hidden]
|
| 1181 |
+
if 'block_outputs' in sliced:
|
| 1182 |
+
for mod in sliced['block_outputs']:
|
| 1183 |
+
out = sliced['block_outputs'][mod]['output']
|
| 1184 |
+
if isinstance(out, list) and len(out) > 0 and isinstance(out[0], list):
|
| 1185 |
+
if pos < len(out[0]):
|
| 1186 |
+
sliced['block_outputs'][mod]['output'] = [[out[0][pos]]]
|
| 1187 |
+
|
| 1188 |
+
# Slice Attention Outputs: [batch, heads, seq, seq] -> [batch, heads, 1, seq]
|
| 1189 |
+
if 'attention_outputs' in sliced:
|
| 1190 |
+
for mod in sliced['attention_outputs']:
|
| 1191 |
+
out = sliced['attention_outputs'][mod]['output']
|
| 1192 |
+
if len(out) > 1:
|
| 1193 |
+
attns = out[1]
|
| 1194 |
+
if isinstance(attns, list) and len(attns) > 0:
|
| 1195 |
+
batch_0 = attns[0]
|
| 1196 |
+
new_batch_0 = []
|
| 1197 |
+
for head in batch_0:
|
| 1198 |
+
if pos < len(head):
|
| 1199 |
+
new_batch_0.append([head[pos]])
|
| 1200 |
+
sliced['attention_outputs'][mod]['output'] = [out[0], [new_batch_0]] + out[2:]
|
| 1201 |
+
|
| 1202 |
+
# Slice input_ids
|
| 1203 |
+
if 'input_ids' in sliced:
|
| 1204 |
+
ids = sliced['input_ids'][0]
|
| 1205 |
+
if pos < len(ids):
|
| 1206 |
+
sliced['input_ids'][0] = ids[:pos+1]
|
| 1207 |
+
|
| 1208 |
+
return sliced
|
| 1209 |
+
|
| 1210 |
+
# Initialize matrix and top_tokens 2D array
|
| 1211 |
+
matrix = [[0.0] * seq_len for _ in range(num_layers)]
|
| 1212 |
+
top_tokens_matrix = [[''] * seq_len for _ in range(num_layers)]
|
| 1213 |
+
|
| 1214 |
+
# Compute delta for each position
|
| 1215 |
+
for pos in range(seq_len):
|
| 1216 |
+
sliced = slice_data(activation_data, pos)
|
| 1217 |
+
layer_data = extract_layer_data(sliced, model, tokenizer)
|
| 1218 |
+
|
| 1219 |
+
if not layer_data:
|
| 1220 |
+
continue
|
| 1221 |
+
|
| 1222 |
+
# Fill in matrix for this position
|
| 1223 |
+
for layer_info_item in layer_data:
|
| 1224 |
+
layer_num = layer_info_item.get('layer_num')
|
| 1225 |
+
if layer_num is None or layer_num not in layer_nums:
|
| 1226 |
+
continue
|
| 1227 |
+
|
| 1228 |
+
layer_idx = layer_nums.index(layer_num)
|
| 1229 |
+
|
| 1230 |
+
# Get top token and its delta (layer-to-layer change)
|
| 1231 |
+
top_token = layer_info_item.get('top_token', '')
|
| 1232 |
+
deltas = layer_info_item.get('deltas', {})
|
| 1233 |
+
|
| 1234 |
+
# The delta for the top token represents how much it changed from prev layer
|
| 1235 |
+
delta = deltas.get(top_token, 0.0) if top_token else 0.0
|
| 1236 |
+
|
| 1237 |
+
matrix[layer_idx][pos] = delta
|
| 1238 |
+
top_tokens_matrix[layer_idx][pos] = top_token if top_token else ''
|
| 1239 |
+
|
| 1240 |
+
return {
|
| 1241 |
+
'matrix': matrix,
|
| 1242 |
+
'tokens': tokens,
|
| 1243 |
+
'layer_nums': layer_nums,
|
| 1244 |
+
'top_tokens': top_tokens_matrix
|
| 1245 |
+
}
|
| 1246 |
+
|
| 1247 |
+
|
| 1248 |
def compute_layer_wise_summaries(layer_data: List[Dict[str, Any]], activation_data: Dict[str, Any]) -> Dict[str, Any]:
|
| 1249 |
"""
|
| 1250 |
Compute summary structures from layer data for easy access.
|