cdpearlman commited on
Commit
ac745db
·
1 Parent(s): e441b65

heatmap implementation

Browse files
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, 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'}, {}, 0, {}, 0, True
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'}, {}, 0, {}, 0, True
589
 
590
  # Run forward pass on the Generated Text
591
  activation_data = execute_forward_pass(model, tokenizer, text, config)
592
 
593
- # Setup scrubber
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'}, {}, 0, {}, 0, True
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 # Or show error
658
 
659
  # Run forward pass
660
  activation_data = execute_forward_pass(model, tokenizer, text, config)
661
 
662
- # Setup scrubber
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
- # Update create_layer_accordions to handle scrubber
675
- # Replaces previous implementation
676
  @app.callback(
677
- Output('layer-accordions-container', 'children'),
 
 
678
  [Input('session-activation-store', 'data'),
679
  Input('session-activation-store-2', 'data'),
680
  Input('session-activation-store-original', 'data'),
681
- Input('sequence-scrubber', 'value')],
682
  [State('model-dropdown', 'value')]
683
  )
684
- def create_layer_accordions(activation_data, activation_data2, original_activation_data, scrubber_val, model_name):
685
- """Create accordion panels for each layer with top-5 bar charts and deltas."""
 
 
 
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 (e.g. from storage quota errors)
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
- # SLICING LOGIC
708
- # We need to slice activation_data to represent the state at step `scrubber_val`
709
- # `scrubber_val` corresponds to the position index.
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
- def slice_data(data, pos):
715
- if not data: return data
716
- sliced = copy.deepcopy(data)
717
-
718
- # Slice Block Outputs: [batch, seq, hidden] -> [batch, 1, hidden]
719
- if 'block_outputs' in sliced:
720
- for mod in sliced['block_outputs']:
721
- # output is [1, seq, hidden] list/tensor
722
- out = sliced['block_outputs'][mod]['output']
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
- activation_data = slice_data(activation_data, scrubber_val)
812
- if activation_data2:
813
- activation_data2 = slice_data(activation_data2, scrubber_val)
814
- if original_activation_data:
815
- original_activation_data = slice_data(original_activation_data, scrubber_val)
816
-
817
-
818
- # Check if we're in ablation mode
819
- ablation_mode = activation_data.get('ablated', False) and original_activation_data
820
-
821
- # Extract layer data for current activation (may be ablated)
822
- layer_data = extract_layer_data(activation_data, model, tokenizer)
823
-
824
- if not layer_data:
825
- return html.P("No layer data available.", className="placeholder-text")
826
-
827
- # Compute layer-wise probability tracking
828
- tracking_data = compute_layer_wise_summaries(layer_data, activation_data)
829
- layer_wise_probs = tracking_data.get('layer_wise_top5_probs', {})
830
- significant_layers = tracking_data.get('significant_layers', [])
831
- # Get global top 5 tokens from activation data
832
- global_top5 = activation_data.get('global_top5_tokens', [])
833
-
834
- # Ensure global_top5 is list of dicts (handle legacy tuples/lists from old sessions)
835
- if global_top5 and isinstance(global_top5[0], (list, tuple)):
836
- global_top5 = [{'token': t, 'probability': p} for t, p in global_top5]
837
-
838
- # If in ablation mode, also extract original layer data
839
- original_layer_data = None
840
- original_layer_wise_probs = {}
841
- original_significant_layers = []
842
- original_global_top5 = []
843
-
844
- if ablation_mode:
845
- original_layer_data = extract_layer_data(original_activation_data, model, tokenizer)
846
- original_tracking_data = compute_layer_wise_summaries(original_layer_data, original_activation_data)
847
- original_layer_wise_probs = original_tracking_data.get('layer_wise_top5_probs', {})
848
- original_significant_layers = original_tracking_data.get('significant_layers', [])
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
- # Create transformer layer flow diagram
1363
- transformer_diagram = html.Div([
1364
- # Input vector
1365
- html.Div([
1366
- html.Div("[ ... ]", className="flow-box", title="This layer receives the output from the previous layer. Each layer builds on what earlier layers learned, gradually understanding the text better."),
1367
- html.Div("Input", style={'fontSize': '11px', 'color': '#6c757d', 'textAlign': 'center'})
1368
- ], style={'display': 'inline-block', 'verticalAlign': 'middle'}),
1369
-
1370
- # Arrow to Self-Attention
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
- # Return all components
1473
- return html.Div([
1474
- layers_container, # Collapsible layers container at top
1475
- *line_graphs, # Line graph(s) below (showing outputs)
1476
- full_bertviz_section # Full BertViz button at the bottom
1477
- ])
 
 
 
 
 
1478
 
1479
  except Exception as e:
1480
- print(f"Error creating accordions: {e}")
1481
  import traceback
1482
  traceback.print_exc()
1483
- return html.P(f"Error creating layer view: {str(e)}", className="placeholder-text")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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', 'paddingRight': '10px'})
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
- # Scrubber
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="layer-accordions-loading",
101
  type="default",
102
- children=html.Div(id='layer-accordions-container', className="layer-accordions"),
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 visuals...", style={'fontSize': '16px', 'color': '#495057'})
107
  ], style={'display': 'flex', 'alignItems': 'center', 'justifyContent': 'center', 'padding': '2rem'})
108
  ),
109
 
110
- # Sequence Ablation Results (New)
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'}) # Hidden by default
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- return html.Div([
102
- html.H4("Your Prompt's Tokenization:",
103
- style={'marginTop': '1.5rem', 'marginBottom': '1rem',
104
- 'color': '#495057', 'fontSize': '16px'}),
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
105
 
106
- # Three-column grid
107
  html.Div([
108
- # Column 1: Tokens
109
- html.Div([
110
- html.H5("Tokens", style={'marginBottom': '1rem', 'color': '#495057'}),
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
- # Column 2: Token IDs
118
  html.Div([
119
- html.H5("Token IDs", style={'marginBottom': '1rem', 'color': '#495057'}),
120
  html.Div([
121
- create_token_box(str(token_id), color, idx, 'id')
122
- for idx, (token_id, color) in enumerate(zip(token_ids_list, color_palette))
123
- ], className='token-column')
124
- ], className='tokenization-col', style={'flex': '1'}),
125
-
126
- # Column 3: Embeddings
127
- html.Div([
128
- html.H5("Embeddings", style={'marginBottom': '1rem', 'color': '#495057'}),
 
 
 
 
 
 
 
 
 
129
  html.Div([
130
- create_token_box("[ ... ]", color, idx, 'embedding')
131
- for idx, color in enumerate(color_palette)
132
- ], className='token-column')
133
- ], className='tokenization-col', style={'flex': '1'})
 
 
 
 
 
134
 
135
- ], className='tokenization-grid',
136
- style={'display': 'flex', 'gap': '2rem', 'alignItems': 'flex-start'})
137
 
138
- ], style={'marginTop': '1rem', 'padding': '1rem', 'backgroundColor': '#ffffff',
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
- - [ ] Add padding/margin to `.glossary-content` in `components/glossary.py`
7
- - [ ] Test that text no longer reaches screen edges
8
 
9
  ### Collapsible Tokenization Example
10
- - [ ] Wrap static diagram in `html.Details()` in `components/tokenization_panel.py`
11
- - [ ] Add `html.Summary("View example tokenization flow")`
12
- - [ ] Default to collapsed state
 
 
 
 
13
 
14
  ### Fix Double Arrow (Transformer Layers)
15
- - [ ] Add `.transformer-layers-summary::-webkit-details-marker { display: none; }` to `assets/style.css`
16
- - [ ] Add `.transformer-layers-summary::marker { display: none; }` for Firefox
17
 
18
  ### Reorder Tokenization Section
19
- - [ ] Move `create_tokenization_panel()` above scrubber container in `components/main_panel.py`
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
- - [ ] Create `compute_position_layer_matrix()` in `utils/model_patterns.py`
32
- - [ ] Reuse existing `slice_data()` logic for per-position slicing
33
- - [ ] Return 2D array: `[num_layers, seq_len]` of delta values
34
 
35
  ### Heatmap Component
36
- - [ ] Create Plotly `go.Heatmap` with `Blues` colorscale
37
- - [ ] Set X-axis labels to token strings
38
- - [ ] Set Y-axis labels to layer numbers (reversed for bottom-up)
39
- - [ ] Add hover template showing token, layer, delta
40
 
41
  ### UI Integration
42
- - [ ] Remove scrubber container from `components/main_panel.py`
43
- - [ ] Add `html.Div(id='heatmap-container')` in its place
44
- - [ ] Create callback to render heatmap from activation data
 
45
 
46
  ### Modal Interaction
47
- - [ ] Create modal component (reuse accordion content structure)
48
- - [ ] Add callback on heatmap `clickData` to extract (layer, position)
49
- - [ ] Populate modal with top-5 chart, attention viz, deltas
50
- - [ ] Add close button and context header
 
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.