a0y0346 commited on
Commit
7e37337
·
1 Parent(s): 7720e5f

Remove Visualizer tab - focus on GPU benchmarks and real model analysis

Browse files
Files changed (1) hide show
  1. app.py +3 -162
app.py CHANGED
@@ -31,12 +31,6 @@ from src.constants import (
31
  SEQ_LENGTH_OPTIONS,
32
  )
33
  from src.models import get_available_models, get_model_memory_footprint
34
- from src.visualizer import (
35
- create_tiling_grid,
36
- create_online_softmax_state,
37
- create_memory_hierarchy_diagram,
38
- get_max_steps,
39
- )
40
  from src.prefill_decode import (
41
  get_attention_pattern_chart,
42
  )
@@ -63,167 +57,14 @@ def create_app() -> gr.Blocks:
63
  # FlashAttention Explorer
64
 
65
  Interactive demonstrations for the **Attention Optimizations** article series.
66
- Explore FlashAttention tiling, online softmax, GQA/MQA, and memory budgets
67
- using real HuggingFace models.
68
 
69
  **Models used:** SmolLM2-360M, Qwen2.5-0.5B, Llama-3.2-1B
70
  """)
71
 
72
  with gr.Tabs():
73
- # Tab 1: Visualizer (CPU)
74
- with gr.Tab("Visualizer", id="tab-visualizer"):
75
- gr.Markdown("""
76
- ## FlashAttention Visualizer
77
-
78
- Understand how FlashAttention processes attention in tiles,
79
- avoiding the O(N²) memory bottleneck. Step through the algorithm
80
- to see how tiles are processed and how online softmax maintains
81
- running statistics.
82
- """)
83
-
84
- # Controls
85
- with gr.Row():
86
- with gr.Column(scale=1):
87
- seq_len_viz = gr.Slider(
88
- minimum=4,
89
- maximum=16,
90
- step=2,
91
- value=8,
92
- label="Sequence Length (tokens)",
93
- )
94
- with gr.Column(scale=1):
95
- block_size_viz = gr.Slider(
96
- minimum=2,
97
- maximum=4,
98
- step=1,
99
- value=2,
100
- label="Block Size",
101
- )
102
- with gr.Column(scale=1):
103
- causal_viz = gr.Checkbox(
104
- value=False,
105
- label="Causal Masking",
106
- )
107
-
108
- # Step controls
109
- with gr.Row():
110
- step_back_btn = gr.Button("◀ Step Back", size="sm")
111
- step_slider = gr.Slider(
112
- minimum=0,
113
- maximum=15,
114
- step=1,
115
- value=0,
116
- label="Current Step",
117
- )
118
- step_forward_btn = gr.Button("Step Forward ▶", size="sm")
119
- reset_btn = gr.Button("Reset", size="sm", variant="secondary")
120
-
121
- # Tiling and Online Softmax side by side
122
- with gr.Row():
123
- with gr.Column(scale=1):
124
- gr.Markdown("### Attention Matrix Tiling")
125
- tiling_plot = gr.Plot(label="Tiling View")
126
-
127
- with gr.Column(scale=1):
128
- gr.Markdown("### Online Softmax State")
129
- softmax_plot = gr.Plot(label="Running m and l")
130
- softmax_explanation = gr.Markdown("*Step through to see online softmax updates*")
131
-
132
- # Memory Hierarchy
133
- gr.Markdown("### Memory Hierarchy Comparison")
134
- with gr.Row():
135
- algo_choice = gr.Radio(
136
- choices=["flash", "standard"],
137
- value="flash",
138
- label="Algorithm",
139
- )
140
- memory_plot = gr.Plot(label="Memory Hierarchy")
141
-
142
- # Event handlers for visualizer
143
- def update_visualizations(seq_len, block_size, causal, step):
144
- """Update all visualizations based on current parameters."""
145
- max_steps = get_max_steps(seq_len, block_size, causal)
146
- # Clamp step to valid range
147
- step = min(step, max_steps - 1)
148
- step = max(step, 0)
149
-
150
- tiling_fig = create_tiling_grid(seq_len, block_size, step, causal)
151
-
152
- # Online softmax uses 4 tiles for the example
153
- num_tiles = seq_len // block_size
154
- softmax_step = min(step, num_tiles - 1)
155
- softmax_fig, explanation = create_online_softmax_state(softmax_step, num_tiles)
156
-
157
- # Only return plots, not step - step is controlled by buttons
158
- return tiling_fig, softmax_fig, explanation
159
-
160
- def update_memory_hierarchy(algo):
161
- """Update memory hierarchy diagram."""
162
- return create_memory_hierarchy_diagram(algo)
163
-
164
- def step_forward(seq_len, block_size, causal, current_step):
165
- """Move to next step."""
166
- max_steps = get_max_steps(seq_len, block_size, causal)
167
- new_step = min(current_step + 1, max_steps - 1)
168
- return new_step
169
-
170
- def step_back(current_step):
171
- """Move to previous step."""
172
- return max(current_step - 1, 0)
173
-
174
- def reset_step(seq_len, block_size, causal):
175
- """Reset to step 0 and update visualizations."""
176
- step = 0
177
- tiling_fig = create_tiling_grid(seq_len, block_size, step, causal)
178
- num_tiles = seq_len // block_size
179
- softmax_fig, explanation = create_online_softmax_state(step, num_tiles)
180
- return step, tiling_fig, softmax_fig, explanation
181
-
182
- # Wire up events
183
- viz_inputs = [seq_len_viz, block_size_viz, causal_viz, step_slider]
184
- viz_outputs = [tiling_plot, softmax_plot, softmax_explanation]
185
-
186
- # Update on parameter change - these don't modify step
187
- seq_len_viz.change(fn=update_visualizations, inputs=viz_inputs, outputs=viz_outputs)
188
- block_size_viz.change(fn=update_visualizations, inputs=viz_inputs, outputs=viz_outputs)
189
- # Causal toggle resets step to 0 since the grid structure changes
190
- causal_viz.change(
191
- fn=reset_step,
192
- inputs=[seq_len_viz, block_size_viz, causal_viz],
193
- outputs=[step_slider, tiling_plot, softmax_plot, softmax_explanation]
194
- )
195
- # Step slider changes only update plots, not the slider itself
196
- step_slider.change(fn=update_visualizations, inputs=viz_inputs, outputs=viz_outputs)
197
-
198
- # Step controls
199
- step_forward_btn.click(
200
- fn=step_forward,
201
- inputs=[seq_len_viz, block_size_viz, causal_viz, step_slider],
202
- outputs=step_slider
203
- )
204
- step_back_btn.click(fn=step_back, inputs=step_slider, outputs=step_slider)
205
- reset_btn.click(
206
- fn=reset_step,
207
- inputs=[seq_len_viz, block_size_viz, causal_viz],
208
- outputs=[step_slider, tiling_plot, softmax_plot, softmax_explanation]
209
- )
210
-
211
- # Memory hierarchy
212
- algo_choice.change(fn=update_memory_hierarchy, inputs=algo_choice, outputs=memory_plot)
213
-
214
- # Initialize on load
215
- demo.load(
216
- fn=update_visualizations,
217
- inputs=viz_inputs,
218
- outputs=viz_outputs
219
- )
220
- demo.load(
221
- fn=update_memory_hierarchy,
222
- inputs=algo_choice,
223
- outputs=memory_plot
224
- )
225
-
226
- # Tab 2: Benchmark (Zero GPU)
227
  with gr.Tab("Benchmark", id="tab-benchmark"):
228
  gr.Markdown("""
229
  ## Attention Backend Benchmark
 
31
  SEQ_LENGTH_OPTIONS,
32
  )
33
  from src.models import get_available_models, get_model_memory_footprint
 
 
 
 
 
 
34
  from src.prefill_decode import (
35
  get_attention_pattern_chart,
36
  )
 
57
  # FlashAttention Explorer
58
 
59
  Interactive demonstrations for the **Attention Optimizations** article series.
60
+ Benchmark attention backends, explore GQA/MQA, prefill vs decode phases,
61
+ and memory budgets using real HuggingFace models on GPU.
62
 
63
  **Models used:** SmolLM2-360M, Qwen2.5-0.5B, Llama-3.2-1B
64
  """)
65
 
66
  with gr.Tabs():
67
+ # Tab 1: Benchmark (Zero GPU)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
68
  with gr.Tab("Benchmark", id="tab-benchmark"):
69
  gr.Markdown("""
70
  ## Attention Backend Benchmark