-
Notifications
You must be signed in to change notification settings - Fork 593
Add KDA #621
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Add KDA #621
Changes from 85 commits
0f6f8ce
3720f19
0307e92
6927239
5e8e1df
aea459d
418694c
ca72ca6
8010348
3637c18
52049ae
f2417a8
56397fd
6ab853d
dcc683f
2ace7b8
45d6354
31be5cb
40b36f9
7b71555
99838ac
98a2cd8
3ead289
dcc8e89
941e39b
5248b29
4fd1e57
5deb9e3
2c50777
7de5dca
7404ac1
9495501
b928910
bc6a1b7
6b54742
80696e4
ae38cd8
4342647
1c482b7
c541794
671a7ea
973266b
224f9de
2236817
3314a79
4485c0e
11ab983
71c9051
003cf3e
7287fce
e23db7a
07ea966
0bfa136
10fba32
2a97b66
d98eebd
be68f24
e9b6eee
800d20c
ba7971b
0550e59
f96d22a
6989914
dea5c64
ddaa1c4
bb54da0
ecb6554
ce6da61
37f454f
65004ee
b58a7dc
ff7bb72
8bf7d25
aa8bb0c
11dca3a
6215fd2
9064809
7c3c8f3
0b45606
0da4bb0
0b11984
669ce78
5774ff8
2b65e00
26fe007
00b1783
082d92b
eccea7e
5b8ed72
0511a50
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| @@ -0,0 +1,149 @@ | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| # -*- coding: utf-8 -*- | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| import os | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| import torch | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| import triton | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| from flash_attn import flash_attn_func | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| from torch.nn import functional as F | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| from fla.ops.comba import chunk_comba | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| from fla.ops.gated_delta_rule import chunk_gated_delta_rule | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| from fla.ops.generalized_delta_rule import chunk_dplr_delta_rule | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| from fla.ops.kda import chunk_kda | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| @triton.testing.perf_report( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| triton.testing.Benchmark( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| # argument names to use as an x-axis for the plot | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| x_names=['T'], | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| # different possible values for `x_name` | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| x_vals=[256, 512, 1024, 2048, 4096, 8192, 16384, 32768, 65536], | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| # argument name whose value corresponds to a different line in the plot | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| line_arg='provider', | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| # possible values for `line_arg`` | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| line_vals=['gdn', 'comba', 'kda', 'dplr'], | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| # label name for the lines | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| line_names=['gdn', 'comba', 'kda', 'dplr'], | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| # line styles | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| styles=[('blue', '-'), ('red', '-.'), ('green', '-'), ('orange', '-.'), | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ('purple', '-'), ('brown', '-.'), ('pink', '-'), ('gray', '-.')], | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ylabel="Execution Time (ms)", # label name for the y-axis | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| # name for the plot. Used also as a file name for saving the plot. | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| plot_name="Performance", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| args={}, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| def benchmark(T, provider): | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| from fla.utils import device | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| dtype = torch.bfloat16 | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| B, H, D = 1, 16, 128 | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| # Set TMA environment variable based on provider | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| original_tma_env = os.environ.get('FLA_USE_TMA', '0') | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| if provider.endswith('_no_tma'): | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| os.environ['FLA_USE_TMA'] = '0' | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| provider_base = provider.replace('_no_tma', '') | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| else: | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| os.environ['FLA_USE_TMA'] = '1' | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| provider_base = provider | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| quantiles = [0.5, 0.2, 0.8] | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| results = 0, 0, 0 | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| do = torch.randn(B, T, H, D, dtype=dtype, device=device) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| if provider_base == 'gdn': | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| q = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| k = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| v = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| g = F.logsigmoid(torch.randn(B, T, H, dtype=dtype, device=device)).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| beta = torch.randn(B, T, H, dtype=dtype, device=device).sigmoid().requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| results = triton.testing.do_bench( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| lambda: chunk_gated_delta_rule( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| q=q, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| k=k, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| v=v, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| g=g, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| beta=beta, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| use_qk_l2norm_in_kernel=True, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| )[0].backward(do), | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| quantiles=quantiles | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| elif provider_base == 'attn': | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| q = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| k = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| v = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| results = triton.testing.do_bench( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| lambda: flash_attn_func( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| q=q, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| k=k, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| v=v, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ).backward(do), | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| quantiles=quantiles | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| if provider_base == 'comba': | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| q = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| k = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| p = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| v = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| g = F.logsigmoid(torch.randn(B, T, H, dtype=torch.float, device=device)).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
zhiyuan1i marked this conversation as resolved.
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| beta = torch.randn(B, T, H, dtype=dtype, device=device).sigmoid().requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| results = triton.testing.do_bench( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| lambda: chunk_comba( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| q=q, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| k=k, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| p=p, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| v=v, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| g=g, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| beta=beta, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| use_qk_l2norm_in_kernel=True, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| )[0].backward(do), | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| quantiles=quantiles | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| if provider_base == 'kda': | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| q = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| k = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| p = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
zhiyuan1i marked this conversation as resolved.
Outdated
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| v = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| g = F.logsigmoid(torch.randn(B, T, H, D, dtype=dtype, device=device)).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| beta = torch.randn(B, T, H, dtype=dtype, device=device).sigmoid().requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| results = triton.testing.do_bench( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| lambda: chunk_kda( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| q=q, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| k=k, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| v=v, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| g=g, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| beta=beta, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| use_qk_l2norm_in_kernel=True, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| )[0].backward(do), | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| quantiles=quantiles | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
Comment on lines
+104
to
+120
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Fix critical shape mismatch for Line 108 creates Apply this diff to fix the shape: - g = F.logsigmoid(torch.randn(B, T, H, D, dtype=dtype, device=device)).requires_grad_(True)
+ g = F.logsigmoid(torch.randn(B, T, H, dtype=dtype, device=device)).requires_grad_(True)📝 Committable suggestion
Suggested change
🤖 Prompt for AI Agents |
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| elif provider_base == 'dplr': | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| q = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| k = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| a = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| b = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| v = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| g = F.logsigmoid(torch.randn(B, T, H, D, dtype=dtype, device=device)).requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| beta = torch.randn(B, T, H, dtype=dtype, device=device).sigmoid().requires_grad_(True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| results = triton.testing.do_bench( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| lambda: chunk_dplr_delta_rule( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| q=q, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| k=k, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| v=v, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| a=a, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| b=b, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| gk=g, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| )[0].backward(do), | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| quantiles=quantiles | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
zhiyuan1i marked this conversation as resolved.
Comment on lines
+121
to
+139
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Remove unused Line 128 creates a Apply this diff to remove the unused variable: v = torch.randn(B, T, H, D, dtype=dtype, device=device).requires_grad_(True)
g = F.logsigmoid(torch.randn(B, T, H, D, dtype=dtype, device=device)).requires_grad_(True)
- beta = torch.randn(B, T, H, dtype=dtype, device=device).sigmoid().requires_grad_(True)
results = triton.testing.do_bench(📝 Committable suggestion
Suggested change
🤖 Prompt for AI Agents |
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| # Restore original TMA environment variable | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| os.environ['FLA_USE_TMA'] = original_tma_env | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| return results | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| if __name__ == '__main__': | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| benchmark.run(print_data=True) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Wrap environment variable restoration in try-finally to ensure cleanup.
The TMA environment variable restoration at line 142 won't execute if an exception occurs during benchmarking. This could leave the environment in an inconsistent state for subsequent runs.
Apply this diff to ensure proper cleanup:
quantiles = [0.5, 0.2, 0.8] results = 0, 0, 0 - do = torch.randn(B, T, H, D, dtype=dtype, device=device) - if provider_base == 'gdn': + try: + do = torch.randn(B, T, H, D, dtype=dtype, device=device) + if provider_base == 'gdn': + ... + elif provider_base == 'dplr': + ... + finally: + # Restore original TMA environment variable + os.environ['FLA_USE_TMA'] = original_tma_env + - # Restore original TMA environment variable - os.environ['FLA_USE_TMA'] = original_tma_env return resultsAlso applies to: 141-142
🤖 Prompt for AI Agents