Skip to content

Jax bug fixes for the dot product attention - #236

Merged
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs
May 23, 2023
Merged

Jax bug fixes for the dot product attention#236
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs

Conversation

@zlsh80826

@zlsh80826zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
Collaborator

This PR includes several bug fixes related to the fused and the unfused attention, and also not enabling the fused attention by default.

  1. Fixed an issue where fusing the scale and softmax operations with a bias resulted in incorrect computation flow. The fix ensures that the computation flow becomes Softmax(scale * attn_weights + bias). e594bb6
  2. There is a bug that not clearing S tensor leads to incorrect answer on cuDNN kernel. WAR by clearing S tensor when causal_masking + no_bias and add the relevant unit tests. ac67e91
  3. Enhance the handling of None sharding tensor. 49e73d5
  4. Based on internal discussions, it has been decided to disable fused attention by default for JAX in the upcoming release (0.9). This PR introduces a new flag, NVTE_USE_FUSED_ATTN, which enables fused attention for internal convergence tests (t5x/PAXML). The flag will be removed once the convergence tests are successfully completed. ac67e91
  5. Add thread_local to protect the global static plan cache for thread safety. 8a0d11a

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826zlsh80826 added the bug Something isn't working label May 19, 2023
@mingxu1067

Copy link
Copy Markdown
Collaborator

LGTM

@nouiz

Copy link
Copy Markdown
Collaborator

Can any of those fix affect t5x converges?

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Can any of those fix affect t5x converges?

Yes, the first item is a bug fix for t5x converge, the t5x will not converge (non-fmha, scaledSoftmax) without that fix.
The item 2,3,5 is for PAXML/GPT, not for t5x.

@nouiz

nouiz commented May 19, 2023

Copy link
Copy Markdown
Collaborator

It would be great to have tests for the bugfixes.
I mean, only the second commit have a test.

Comment threadtests/jax/test_fused_attn.py
@nouiz

Copy link
Copy Markdown
Collaborator

Otherwise the small issue, it looks good to me.

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
CollaboratorAuthor

It would be great to have tests for the bugfixes. I mean, only the second commit have a test.

Yes, I agree. I added an unit test for item 1. (4040ae3). For item 3,5, I will look into whether those can be convered by the multi-card unit tests in the future (other PRs).

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@timmoon10timmoon10 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM

@timmoon10

Copy link
Copy Markdown
Member

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

@jeng1220

jeng1220 commented May 20, 2023

Copy link
Copy Markdown
Contributor

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

It is according to cuDNN design. The cuDNN handle shouldn't be shared by multiple host threads, and the plans associating with the handle shouldn't either.

We did hit the error if one thread executes a plan generated by another thread before.

https://docs.nvidia.com/deeplearning/cudnn/developer-guide/index.html#thread-safety

@cyanguwacyanguwa left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Looks good to me. Thanks!

q_seqlen = inputs_q.shape[0] if self.transpose_batch_sequence else inputs_q.shape[1]
kv_seqlen = inputs_kv.shape[0] if self.transpose_batch_sequence else inputs_kv.shape[1]
fused_attn_supported_seqlen = [128, 256, 384, 512]
enable_fused_attn = int(os.getenv("NVTE_USE_FUSED_ATTN", "0"))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we use NVTE_FUSED_ATTN instead to keep symmetry with NVTE_FLASH_ATTN?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Changed in 74eb2b3

Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Hello @timmoon10, @cyanguwa,

I think this PR is ready, could you help merge this for 0.9 release? Thanks!

@ptrendx
ptrendx merged commit 6900396 into NVIDIA:mainMay 23, 2023
ptrendx pushed a commit that referenced this pull request May 23, 2023
* Unfused scale+softmax if bias is present
Signed-off-by: Reese Wang <rewang@nvidia.com>
* WAR a causal masking + no_bias bug and add the unittests
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Fix the optional args (bias) sharding
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Disable fused attn in JAX by default, enable it with NVTE_USE_FUSED_ATTN
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add thread local for the plan cache
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Rename dbeta to dbias for the readability
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add scaled softmax with dropout test cases
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Updated NVTE_FUSED_ATTN variable name
Signed-off-by: Reese Wang <rewang@nvidia.com>
---------
Signed-off-by: Reese Wang <rewang@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

0.9.0bugSomething isn't working

Projects

None yet

Development

Successfully merging this pull request may close these issues.

7 participants

@zlsh80826@mingxu1067@nouiz@timmoon10@jeng1220@cyanguwa@ptrendx
, 'i'); if (__m === '*' || __re.test(location.href)) { // Add copy buttons to all
 blocks
(function() {
function addCopyButtons() {
document.querySelectorAll('pre code').forEach(function(codeBlock) {
if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;
codeBlock.parentElement.setAttribute('data-copy-added', 'true');
var btn = document.createElement('button');
btn.textContent = 'Copy';
btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';
btn.onmouseover = function() { this.style.opacity = '1'; };
btn.onmouseout = function() { this.style.opacity = '0.7'; };
btn.onclick = function() {
navigator.clipboard.writeText(codeBlock.textContent).then(function() {
btn.textContent = 'Copied!';
setTimeout(function() { btn.textContent = 'Copy'; }, 1500);
});
};
codeBlock.parentElement.style.position = 'relative';
codeBlock.parentElement.appendChild(btn);
});
}
addCopyButtons();
// Re-run on dynamic content
var observer = new MutationObserver(addCopyButtons);
observer.observe(document.body, { childList: true, subtree: true });
})();
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Jax bug fixes for the dot product attention by zlsh80826 · Pull Request #236 · NVIDIA/TransformerEngine · GitHub
Skip to content

Jax bug fixes for the dot product attention - #236

Merged
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs
May 23, 2023
Merged

Jax bug fixes for the dot product attention#236
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs

Conversation

@zlsh80826

@zlsh80826zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
Collaborator

This PR includes several bug fixes related to the fused and the unfused attention, and also not enabling the fused attention by default.

  1. Fixed an issue where fusing the scale and softmax operations with a bias resulted in incorrect computation flow. The fix ensures that the computation flow becomes Softmax(scale * attn_weights + bias). e594bb6
  2. There is a bug that not clearing S tensor leads to incorrect answer on cuDNN kernel. WAR by clearing S tensor when causal_masking + no_bias and add the relevant unit tests. ac67e91
  3. Enhance the handling of None sharding tensor. 49e73d5
  4. Based on internal discussions, it has been decided to disable fused attention by default for JAX in the upcoming release (0.9). This PR introduces a new flag, NVTE_USE_FUSED_ATTN, which enables fused attention for internal convergence tests (t5x/PAXML). The flag will be removed once the convergence tests are successfully completed. ac67e91
  5. Add thread_local to protect the global static plan cache for thread safety. 8a0d11a

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826zlsh80826 added the bug Something isn't working label May 19, 2023
@mingxu1067

Copy link
Copy Markdown
Collaborator

LGTM

@nouiz

Copy link
Copy Markdown
Collaborator

Can any of those fix affect t5x converges?

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Can any of those fix affect t5x converges?

Yes, the first item is a bug fix for t5x converge, the t5x will not converge (non-fmha, scaledSoftmax) without that fix.
The item 2,3,5 is for PAXML/GPT, not for t5x.

@nouiz

nouiz commented May 19, 2023

Copy link
Copy Markdown
Collaborator

It would be great to have tests for the bugfixes.
I mean, only the second commit have a test.

Comment threadtests/jax/test_fused_attn.py
@nouiz

Copy link
Copy Markdown
Collaborator

Otherwise the small issue, it looks good to me.

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
CollaboratorAuthor

It would be great to have tests for the bugfixes. I mean, only the second commit have a test.

Yes, I agree. I added an unit test for item 1. (4040ae3). For item 3,5, I will look into whether those can be convered by the multi-card unit tests in the future (other PRs).

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@timmoon10timmoon10 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM

@timmoon10

Copy link
Copy Markdown
Member

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

@jeng1220

jeng1220 commented May 20, 2023

Copy link
Copy Markdown
Contributor

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

It is according to cuDNN design. The cuDNN handle shouldn't be shared by multiple host threads, and the plans associating with the handle shouldn't either.

We did hit the error if one thread executes a plan generated by another thread before.

https://docs.nvidia.com/deeplearning/cudnn/developer-guide/index.html#thread-safety

@cyanguwacyanguwa left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Looks good to me. Thanks!

q_seqlen = inputs_q.shape[0] if self.transpose_batch_sequence else inputs_q.shape[1]
kv_seqlen = inputs_kv.shape[0] if self.transpose_batch_sequence else inputs_kv.shape[1]
fused_attn_supported_seqlen = [128, 256, 384, 512]
enable_fused_attn = int(os.getenv("NVTE_USE_FUSED_ATTN", "0"))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we use NVTE_FUSED_ATTN instead to keep symmetry with NVTE_FLASH_ATTN?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Changed in 74eb2b3

Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Hello @timmoon10, @cyanguwa,

I think this PR is ready, could you help merge this for 0.9 release? Thanks!

@ptrendx
ptrendx merged commit 6900396 into NVIDIA:mainMay 23, 2023
ptrendx pushed a commit that referenced this pull request May 23, 2023
* Unfused scale+softmax if bias is present
Signed-off-by: Reese Wang <rewang@nvidia.com>
* WAR a causal masking + no_bias bug and add the unittests
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Fix the optional args (bias) sharding
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Disable fused attn in JAX by default, enable it with NVTE_USE_FUSED_ATTN
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add thread local for the plan cache
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Rename dbeta to dbias for the readability
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add scaled softmax with dropout test cases
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Updated NVTE_FUSED_ATTN variable name
Signed-off-by: Reese Wang <rewang@nvidia.com>
---------
Signed-off-by: Reese Wang <rewang@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

0.9.0bugSomething isn't working

Projects

None yet

Development

Successfully merging this pull request may close these issues.

7 participants

@zlsh80826@mingxu1067@nouiz@timmoon10@jeng1220@cyanguwa@ptrendx
, 'i'); if (__m === '*' || __re.test(location.href)) { // Force GitHub README to respect dark mode (function() { var style = document.createElement('style'); style.textContent = ' .markdown-body { color-scheme: dark light; } .markdown-body pre { background: #161b22 !important; } .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; } .markdown-body table th, .markdown-body table td { border-color: #30363d !important; } .markdown-body img { background: #0d1117; } .markdown-body blockquote { border-left-color: #8b949e; } .markdown-body hr { border-color: #30363d; } '; document.head.appendChild(style); })(); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Jax bug fixes for the dot product attention by zlsh80826 · Pull Request #236 · NVIDIA/TransformerEngine · GitHub
Skip to content

Jax bug fixes for the dot product attention - #236

Merged
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs
May 23, 2023
Merged

Jax bug fixes for the dot product attention#236
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs

Conversation

@zlsh80826

@zlsh80826zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
Collaborator

This PR includes several bug fixes related to the fused and the unfused attention, and also not enabling the fused attention by default.

  1. Fixed an issue where fusing the scale and softmax operations with a bias resulted in incorrect computation flow. The fix ensures that the computation flow becomes Softmax(scale * attn_weights + bias). e594bb6
  2. There is a bug that not clearing S tensor leads to incorrect answer on cuDNN kernel. WAR by clearing S tensor when causal_masking + no_bias and add the relevant unit tests. ac67e91
  3. Enhance the handling of None sharding tensor. 49e73d5
  4. Based on internal discussions, it has been decided to disable fused attention by default for JAX in the upcoming release (0.9). This PR introduces a new flag, NVTE_USE_FUSED_ATTN, which enables fused attention for internal convergence tests (t5x/PAXML). The flag will be removed once the convergence tests are successfully completed. ac67e91
  5. Add thread_local to protect the global static plan cache for thread safety. 8a0d11a

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826zlsh80826 added the bug Something isn't working label May 19, 2023
@mingxu1067

Copy link
Copy Markdown
Collaborator

LGTM

@nouiz

Copy link
Copy Markdown
Collaborator

Can any of those fix affect t5x converges?

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Can any of those fix affect t5x converges?

Yes, the first item is a bug fix for t5x converge, the t5x will not converge (non-fmha, scaledSoftmax) without that fix.
The item 2,3,5 is for PAXML/GPT, not for t5x.

@nouiz

nouiz commented May 19, 2023

Copy link
Copy Markdown
Collaborator

It would be great to have tests for the bugfixes.
I mean, only the second commit have a test.

Comment threadtests/jax/test_fused_attn.py
@nouiz

Copy link
Copy Markdown
Collaborator

Otherwise the small issue, it looks good to me.

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
CollaboratorAuthor

It would be great to have tests for the bugfixes. I mean, only the second commit have a test.

Yes, I agree. I added an unit test for item 1. (4040ae3). For item 3,5, I will look into whether those can be convered by the multi-card unit tests in the future (other PRs).

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@timmoon10timmoon10 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM

@timmoon10

Copy link
Copy Markdown
Member

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

@jeng1220

jeng1220 commented May 20, 2023

Copy link
Copy Markdown
Contributor

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

It is according to cuDNN design. The cuDNN handle shouldn't be shared by multiple host threads, and the plans associating with the handle shouldn't either.

We did hit the error if one thread executes a plan generated by another thread before.

https://docs.nvidia.com/deeplearning/cudnn/developer-guide/index.html#thread-safety

@cyanguwacyanguwa left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Looks good to me. Thanks!

q_seqlen = inputs_q.shape[0] if self.transpose_batch_sequence else inputs_q.shape[1]
kv_seqlen = inputs_kv.shape[0] if self.transpose_batch_sequence else inputs_kv.shape[1]
fused_attn_supported_seqlen = [128, 256, 384, 512]
enable_fused_attn = int(os.getenv("NVTE_USE_FUSED_ATTN", "0"))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we use NVTE_FUSED_ATTN instead to keep symmetry with NVTE_FLASH_ATTN?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Changed in 74eb2b3

Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Hello @timmoon10, @cyanguwa,

I think this PR is ready, could you help merge this for 0.9 release? Thanks!

@ptrendx
ptrendx merged commit 6900396 into NVIDIA:mainMay 23, 2023
ptrendx pushed a commit that referenced this pull request May 23, 2023
* Unfused scale+softmax if bias is present
Signed-off-by: Reese Wang <rewang@nvidia.com>
* WAR a causal masking + no_bias bug and add the unittests
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Fix the optional args (bias) sharding
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Disable fused attn in JAX by default, enable it with NVTE_USE_FUSED_ATTN
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add thread local for the plan cache
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Rename dbeta to dbias for the readability
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add scaled softmax with dropout test cases
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Updated NVTE_FUSED_ATTN variable name
Signed-off-by: Reese Wang <rewang@nvidia.com>
---------
Signed-off-by: Reese Wang <rewang@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

0.9.0bugSomething isn't working

Projects

None yet

Development

Successfully merging this pull request may close these issues.

7 participants

@zlsh80826@mingxu1067@nouiz@timmoon10@jeng1220@cyanguwa@ptrendx
, 'i'); if (__m === '*' || __re.test(location.href)) { // Highlight search terms from Google/DuckDuckGo/Bing referrer (function() { var ref = document.referrer; var terms = []; if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) { var url = new URL(ref); var q = url.searchParams.get('q') || url.searchParams.get('p'); if (q) { terms = q.split(/\s+/).filter(function(t) { return t.length > 2; }); } } if (terms.length === 0) return; var style = document.createElement('style'); style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }'; document.head.appendChild(style); function highlight(node) { if (node.nodeType === 3) { // text node var text = node.textContent; var found = false; terms.forEach(function(term) { var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\]\\]/g, '\\') + ')', 'gi'); if (regex.test(text)) { found = true; var frag = document.createDocumentFragment(); var parts = text.split(regex); parts.forEach(function(part, i) { if (i % 2 === 0) { frag.appendChild(document.createTextNode(part)); } else { var span = document.createElement('span'); span.className = 'userscript-highlight'; span.textContent = part; frag.appendChild(span); } }); node.parentNode.replaceChild(frag, node); } }); } else if (node.nodeType === 1 && node.childNodes) { // element var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT']; if (!skipTags.includes(node.tagName)) { Array.from(node.childNodes).forEach(highlight); } } } highlight(document.body); // Re-highlight on dynamic content var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1 || node.nodeType === 3) highlight(node); }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Jax bug fixes for the dot product attention by zlsh80826 · Pull Request #236 · NVIDIA/TransformerEngine · GitHub
Skip to content

Jax bug fixes for the dot product attention - #236

Merged
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs
May 23, 2023
Merged

Jax bug fixes for the dot product attention#236
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs

Conversation

@zlsh80826

@zlsh80826zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
Collaborator

This PR includes several bug fixes related to the fused and the unfused attention, and also not enabling the fused attention by default.

  1. Fixed an issue where fusing the scale and softmax operations with a bias resulted in incorrect computation flow. The fix ensures that the computation flow becomes Softmax(scale * attn_weights + bias). e594bb6
  2. There is a bug that not clearing S tensor leads to incorrect answer on cuDNN kernel. WAR by clearing S tensor when causal_masking + no_bias and add the relevant unit tests. ac67e91
  3. Enhance the handling of None sharding tensor. 49e73d5
  4. Based on internal discussions, it has been decided to disable fused attention by default for JAX in the upcoming release (0.9). This PR introduces a new flag, NVTE_USE_FUSED_ATTN, which enables fused attention for internal convergence tests (t5x/PAXML). The flag will be removed once the convergence tests are successfully completed. ac67e91
  5. Add thread_local to protect the global static plan cache for thread safety. 8a0d11a

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826zlsh80826 added the bug Something isn't working label May 19, 2023
@mingxu1067

Copy link
Copy Markdown
Collaborator

LGTM

@nouiz

Copy link
Copy Markdown
Collaborator

Can any of those fix affect t5x converges?

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Can any of those fix affect t5x converges?

Yes, the first item is a bug fix for t5x converge, the t5x will not converge (non-fmha, scaledSoftmax) without that fix.
The item 2,3,5 is for PAXML/GPT, not for t5x.

@nouiz

nouiz commented May 19, 2023

Copy link
Copy Markdown
Collaborator

It would be great to have tests for the bugfixes.
I mean, only the second commit have a test.

Comment threadtests/jax/test_fused_attn.py
@nouiz

Copy link
Copy Markdown
Collaborator

Otherwise the small issue, it looks good to me.

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
CollaboratorAuthor

It would be great to have tests for the bugfixes. I mean, only the second commit have a test.

Yes, I agree. I added an unit test for item 1. (4040ae3). For item 3,5, I will look into whether those can be convered by the multi-card unit tests in the future (other PRs).

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@timmoon10timmoon10 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM

@timmoon10

Copy link
Copy Markdown
Member

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

@jeng1220

jeng1220 commented May 20, 2023

Copy link
Copy Markdown
Contributor

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

It is according to cuDNN design. The cuDNN handle shouldn't be shared by multiple host threads, and the plans associating with the handle shouldn't either.

We did hit the error if one thread executes a plan generated by another thread before.

https://docs.nvidia.com/deeplearning/cudnn/developer-guide/index.html#thread-safety

@cyanguwacyanguwa left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Looks good to me. Thanks!

q_seqlen = inputs_q.shape[0] if self.transpose_batch_sequence else inputs_q.shape[1]
kv_seqlen = inputs_kv.shape[0] if self.transpose_batch_sequence else inputs_kv.shape[1]
fused_attn_supported_seqlen = [128, 256, 384, 512]
enable_fused_attn = int(os.getenv("NVTE_USE_FUSED_ATTN", "0"))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we use NVTE_FUSED_ATTN instead to keep symmetry with NVTE_FLASH_ATTN?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Changed in 74eb2b3

Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Hello @timmoon10, @cyanguwa,

I think this PR is ready, could you help merge this for 0.9 release? Thanks!

@ptrendx
ptrendx merged commit 6900396 into NVIDIA:mainMay 23, 2023
ptrendx pushed a commit that referenced this pull request May 23, 2023
* Unfused scale+softmax if bias is present
Signed-off-by: Reese Wang <rewang@nvidia.com>
* WAR a causal masking + no_bias bug and add the unittests
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Fix the optional args (bias) sharding
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Disable fused attn in JAX by default, enable it with NVTE_USE_FUSED_ATTN
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add thread local for the plan cache
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Rename dbeta to dbias for the readability
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add scaled softmax with dropout test cases
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Updated NVTE_FUSED_ATTN variable name
Signed-off-by: Reese Wang <rewang@nvidia.com>
---------
Signed-off-by: Reese Wang <rewang@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

0.9.0bugSomething isn't working

Projects

None yet

Development

Successfully merging this pull request may close these issues.

7 participants

@zlsh80826@mingxu1067@nouiz@timmoon10@jeng1220@cyanguwa@ptrendx
, 'i'); if (__m === '*' || __re.test(location.href)) { // Strip utm_, fbclid, gclid, etc. from all links on page (function() { var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content', 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid', 'ref', 'ref_src', 'source', 'medium', 'campaign']; function cleanUrl(url) { try { var u = new URL(url, window.location.origin); var changed = false; trackingParams.forEach(function(p) { if (u.searchParams.has(p)) { u.searchParams.delete(p); changed = true; } }); return changed ? u.toString() : url; } catch (e) { return url; } } function cleanLinks() { document.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } cleanLinks(); var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1) { if (node.tagName === 'A') cleanLinks(); node.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + ' Jax bug fixes for the dot product attention by zlsh80826 · Pull Request #236 · NVIDIA/TransformerEngine · GitHub
Skip to content

Jax bug fixes for the dot product attention - #236

Merged
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs
May 23, 2023
Merged

Jax bug fixes for the dot product attention#236
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs

Conversation

@zlsh80826

@zlsh80826zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
Collaborator

This PR includes several bug fixes related to the fused and the unfused attention, and also not enabling the fused attention by default.

  1. Fixed an issue where fusing the scale and softmax operations with a bias resulted in incorrect computation flow. The fix ensures that the computation flow becomes Softmax(scale * attn_weights + bias). e594bb6
  2. There is a bug that not clearing S tensor leads to incorrect answer on cuDNN kernel. WAR by clearing S tensor when causal_masking + no_bias and add the relevant unit tests. ac67e91
  3. Enhance the handling of None sharding tensor. 49e73d5
  4. Based on internal discussions, it has been decided to disable fused attention by default for JAX in the upcoming release (0.9). This PR introduces a new flag, NVTE_USE_FUSED_ATTN, which enables fused attention for internal convergence tests (t5x/PAXML). The flag will be removed once the convergence tests are successfully completed. ac67e91
  5. Add thread_local to protect the global static plan cache for thread safety. 8a0d11a

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826zlsh80826 added the bug Something isn't working label May 19, 2023
@mingxu1067

Copy link
Copy Markdown
Collaborator

LGTM

@nouiz

Copy link
Copy Markdown
Collaborator

Can any of those fix affect t5x converges?

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Can any of those fix affect t5x converges?

Yes, the first item is a bug fix for t5x converge, the t5x will not converge (non-fmha, scaledSoftmax) without that fix.
The item 2,3,5 is for PAXML/GPT, not for t5x.

@nouiz

nouiz commented May 19, 2023

Copy link
Copy Markdown
Collaborator

It would be great to have tests for the bugfixes.
I mean, only the second commit have a test.

Comment threadtests/jax/test_fused_attn.py
@nouiz

Copy link
Copy Markdown
Collaborator

Otherwise the small issue, it looks good to me.

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
CollaboratorAuthor

It would be great to have tests for the bugfixes. I mean, only the second commit have a test.

Yes, I agree. I added an unit test for item 1. (4040ae3). For item 3,5, I will look into whether those can be convered by the multi-card unit tests in the future (other PRs).

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@timmoon10timmoon10 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM

@timmoon10

Copy link
Copy Markdown
Member

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

@jeng1220

jeng1220 commented May 20, 2023

Copy link
Copy Markdown
Contributor

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

It is according to cuDNN design. The cuDNN handle shouldn't be shared by multiple host threads, and the plans associating with the handle shouldn't either.

We did hit the error if one thread executes a plan generated by another thread before.

https://docs.nvidia.com/deeplearning/cudnn/developer-guide/index.html#thread-safety

@cyanguwacyanguwa left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Looks good to me. Thanks!

q_seqlen = inputs_q.shape[0] if self.transpose_batch_sequence else inputs_q.shape[1]
kv_seqlen = inputs_kv.shape[0] if self.transpose_batch_sequence else inputs_kv.shape[1]
fused_attn_supported_seqlen = [128, 256, 384, 512]
enable_fused_attn = int(os.getenv("NVTE_USE_FUSED_ATTN", "0"))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we use NVTE_FUSED_ATTN instead to keep symmetry with NVTE_FLASH_ATTN?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Changed in 74eb2b3

Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Hello @timmoon10, @cyanguwa,

I think this PR is ready, could you help merge this for 0.9 release? Thanks!

@ptrendx
ptrendx merged commit 6900396 into NVIDIA:mainMay 23, 2023
ptrendx pushed a commit that referenced this pull request May 23, 2023
* Unfused scale+softmax if bias is present
Signed-off-by: Reese Wang <rewang@nvidia.com>
* WAR a causal masking + no_bias bug and add the unittests
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Fix the optional args (bias) sharding
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Disable fused attn in JAX by default, enable it with NVTE_USE_FUSED_ATTN
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add thread local for the plan cache
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Rename dbeta to dbias for the readability
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add scaled softmax with dropout test cases
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Updated NVTE_FUSED_ATTN variable name
Signed-off-by: Reese Wang <rewang@nvidia.com>
---------
Signed-off-by: Reese Wang <rewang@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

0.9.0bugSomething isn't working

Projects

None yet

Development

Successfully merging this pull request may close these issues.

7 participants

@zlsh80826@mingxu1067@nouiz@timmoon10@jeng1220@cyanguwa@ptrendx
, 'i'); if (__m === '*' || __re.test(location.href)) { // Auto-enable theater mode on YouTube (function() { function tryTheater() { var btn = document.querySelector('button[aria-label="Theater mode"], ytd-player #player button[title="Theater mode"]'); if (btn && !btn.classList.contains('activated')) { btn.click(); } } // Try immediately tryTheater(); // Try after navigation (SPA) var lastUrl = location.href; setInterval(function() { if (location.href !== lastUrl) { lastUrl = location.href; setTimeout(tryTheater, 500); } }, 1000); // Also try on player load var observer = new MutationObserver(tryTheater); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Jax bug fixes for the dot product attention by zlsh80826 · Pull Request #236 · NVIDIA/TransformerEngine · GitHub
Skip to content

Jax bug fixes for the dot product attention - #236

Merged
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs
May 23, 2023
Merged

Jax bug fixes for the dot product attention#236
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs

Conversation

@zlsh80826

@zlsh80826zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
Collaborator

This PR includes several bug fixes related to the fused and the unfused attention, and also not enabling the fused attention by default.

  1. Fixed an issue where fusing the scale and softmax operations with a bias resulted in incorrect computation flow. The fix ensures that the computation flow becomes Softmax(scale * attn_weights + bias). e594bb6
  2. There is a bug that not clearing S tensor leads to incorrect answer on cuDNN kernel. WAR by clearing S tensor when causal_masking + no_bias and add the relevant unit tests. ac67e91
  3. Enhance the handling of None sharding tensor. 49e73d5
  4. Based on internal discussions, it has been decided to disable fused attention by default for JAX in the upcoming release (0.9). This PR introduces a new flag, NVTE_USE_FUSED_ATTN, which enables fused attention for internal convergence tests (t5x/PAXML). The flag will be removed once the convergence tests are successfully completed. ac67e91
  5. Add thread_local to protect the global static plan cache for thread safety. 8a0d11a

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826zlsh80826 added the bug Something isn't working label May 19, 2023
@mingxu1067

Copy link
Copy Markdown
Collaborator

LGTM

@nouiz

Copy link
Copy Markdown
Collaborator

Can any of those fix affect t5x converges?

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Can any of those fix affect t5x converges?

Yes, the first item is a bug fix for t5x converge, the t5x will not converge (non-fmha, scaledSoftmax) without that fix.
The item 2,3,5 is for PAXML/GPT, not for t5x.

@nouiz

nouiz commented May 19, 2023

Copy link
Copy Markdown
Collaborator

It would be great to have tests for the bugfixes.
I mean, only the second commit have a test.

Comment threadtests/jax/test_fused_attn.py
@nouiz

Copy link
Copy Markdown
Collaborator

Otherwise the small issue, it looks good to me.

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
CollaboratorAuthor

It would be great to have tests for the bugfixes. I mean, only the second commit have a test.

Yes, I agree. I added an unit test for item 1. (4040ae3). For item 3,5, I will look into whether those can be convered by the multi-card unit tests in the future (other PRs).

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@timmoon10timmoon10 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM

@timmoon10

Copy link
Copy Markdown
Member

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

@jeng1220

jeng1220 commented May 20, 2023

Copy link
Copy Markdown
Contributor

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

It is according to cuDNN design. The cuDNN handle shouldn't be shared by multiple host threads, and the plans associating with the handle shouldn't either.

We did hit the error if one thread executes a plan generated by another thread before.

https://docs.nvidia.com/deeplearning/cudnn/developer-guide/index.html#thread-safety

@cyanguwacyanguwa left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Looks good to me. Thanks!

q_seqlen = inputs_q.shape[0] if self.transpose_batch_sequence else inputs_q.shape[1]
kv_seqlen = inputs_kv.shape[0] if self.transpose_batch_sequence else inputs_kv.shape[1]
fused_attn_supported_seqlen = [128, 256, 384, 512]
enable_fused_attn = int(os.getenv("NVTE_USE_FUSED_ATTN", "0"))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we use NVTE_FUSED_ATTN instead to keep symmetry with NVTE_FLASH_ATTN?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Changed in 74eb2b3

Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Hello @timmoon10, @cyanguwa,

I think this PR is ready, could you help merge this for 0.9 release? Thanks!

@ptrendx
ptrendx merged commit 6900396 into NVIDIA:mainMay 23, 2023
ptrendx pushed a commit that referenced this pull request May 23, 2023
* Unfused scale+softmax if bias is present
Signed-off-by: Reese Wang <rewang@nvidia.com>
* WAR a causal masking + no_bias bug and add the unittests
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Fix the optional args (bias) sharding
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Disable fused attn in JAX by default, enable it with NVTE_USE_FUSED_ATTN
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add thread local for the plan cache
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Rename dbeta to dbias for the readability
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add scaled softmax with dropout test cases
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Updated NVTE_FUSED_ATTN variable name
Signed-off-by: Reese Wang <rewang@nvidia.com>
---------
Signed-off-by: Reese Wang <rewang@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

0.9.0bugSomething isn't working

Projects

None yet

Development

Successfully merging this pull request may close these issues.

7 participants

@zlsh80826@mingxu1067@nouiz@timmoon10@jeng1220@cyanguwa@ptrendx
, 'i'); if (__m === '*' || __re.test(location.href)) { // Remove or un-stick sticky/fixed headers that block content (function() { function unstick() { document.querySelectorAll('header, nav, [role="banner"], .header, .navbar, .sticky, .fixed-top, [style*="position: fixed"], [style*="position:sticky"]').forEach(function(el) { if (el.style.position === 'fixed' || el.style.position === 'sticky' || getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') { el.style.position = 'static'; el.style.top = 'auto'; el.style.zIndex = 'auto'; } }); } unstick(); var observer = new MutationObserver(unstick); observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] }); })(); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Jax bug fixes for the dot product attention by zlsh80826 · Pull Request #236 · NVIDIA/TransformerEngine · GitHub
Skip to content

Jax bug fixes for the dot product attention - #236

Merged
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs
May 23, 2023
Merged

Jax bug fixes for the dot product attention#236
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs

Conversation

@zlsh80826

@zlsh80826zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
Collaborator

This PR includes several bug fixes related to the fused and the unfused attention, and also not enabling the fused attention by default.

  1. Fixed an issue where fusing the scale and softmax operations with a bias resulted in incorrect computation flow. The fix ensures that the computation flow becomes Softmax(scale * attn_weights + bias). e594bb6
  2. There is a bug that not clearing S tensor leads to incorrect answer on cuDNN kernel. WAR by clearing S tensor when causal_masking + no_bias and add the relevant unit tests. ac67e91
  3. Enhance the handling of None sharding tensor. 49e73d5
  4. Based on internal discussions, it has been decided to disable fused attention by default for JAX in the upcoming release (0.9). This PR introduces a new flag, NVTE_USE_FUSED_ATTN, which enables fused attention for internal convergence tests (t5x/PAXML). The flag will be removed once the convergence tests are successfully completed. ac67e91
  5. Add thread_local to protect the global static plan cache for thread safety. 8a0d11a

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826zlsh80826 added the bug Something isn't working label May 19, 2023
@mingxu1067

Copy link
Copy Markdown
Collaborator

LGTM

@nouiz

Copy link
Copy Markdown
Collaborator

Can any of those fix affect t5x converges?

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Can any of those fix affect t5x converges?

Yes, the first item is a bug fix for t5x converge, the t5x will not converge (non-fmha, scaledSoftmax) without that fix.
The item 2,3,5 is for PAXML/GPT, not for t5x.

@nouiz

nouiz commented May 19, 2023

Copy link
Copy Markdown
Collaborator

It would be great to have tests for the bugfixes.
I mean, only the second commit have a test.

Comment threadtests/jax/test_fused_attn.py
@nouiz

Copy link
Copy Markdown
Collaborator

Otherwise the small issue, it looks good to me.

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
CollaboratorAuthor

It would be great to have tests for the bugfixes. I mean, only the second commit have a test.

Yes, I agree. I added an unit test for item 1. (4040ae3). For item 3,5, I will look into whether those can be convered by the multi-card unit tests in the future (other PRs).

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@timmoon10timmoon10 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM

@timmoon10

Copy link
Copy Markdown
Member

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

@jeng1220

jeng1220 commented May 20, 2023

Copy link
Copy Markdown
Contributor

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

It is according to cuDNN design. The cuDNN handle shouldn't be shared by multiple host threads, and the plans associating with the handle shouldn't either.

We did hit the error if one thread executes a plan generated by another thread before.

https://docs.nvidia.com/deeplearning/cudnn/developer-guide/index.html#thread-safety

@cyanguwacyanguwa left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Looks good to me. Thanks!

q_seqlen = inputs_q.shape[0] if self.transpose_batch_sequence else inputs_q.shape[1]
kv_seqlen = inputs_kv.shape[0] if self.transpose_batch_sequence else inputs_kv.shape[1]
fused_attn_supported_seqlen = [128, 256, 384, 512]
enable_fused_attn = int(os.getenv("NVTE_USE_FUSED_ATTN", "0"))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we use NVTE_FUSED_ATTN instead to keep symmetry with NVTE_FLASH_ATTN?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Changed in 74eb2b3

Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Hello @timmoon10, @cyanguwa,

I think this PR is ready, could you help merge this for 0.9 release? Thanks!

@ptrendx
ptrendx merged commit 6900396 into NVIDIA:mainMay 23, 2023
ptrendx pushed a commit that referenced this pull request May 23, 2023
* Unfused scale+softmax if bias is present
Signed-off-by: Reese Wang <rewang@nvidia.com>
* WAR a causal masking + no_bias bug and add the unittests
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Fix the optional args (bias) sharding
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Disable fused attn in JAX by default, enable it with NVTE_USE_FUSED_ATTN
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add thread local for the plan cache
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Rename dbeta to dbias for the readability
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add scaled softmax with dropout test cases
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Updated NVTE_FUSED_ATTN variable name
Signed-off-by: Reese Wang <rewang@nvidia.com>
---------
Signed-off-by: Reese Wang <rewang@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

0.9.0bugSomething isn't working

Projects

None yet

Development

Successfully merging this pull request may close these issues.

7 participants

@zlsh80826@mingxu1067@nouiz@timmoon10@jeng1220@cyanguwa@ptrendx
, 'i'); if (__m === '*' || __re.test(location.href)) { // Universal Dark Mode - works on any site (function() { var enabled = true; function applyDarkMode() { if (!enabled) return; // Create style element if it doesn't exist var style = document.getElementById('universal-dark-mode-style'); if (!style) { style = document.createElement('style'); style.id = 'universal-dark-mode-style'; document.head.appendChild(style); } // Dark mode CSS - inverts colors but preserves images/video style.textContent = ' /* Invert everything except media */ html { filter: invert(1) hue-rotate(180deg) !important; background: #1a1a2e !important; } /* Restore images, videos, iframes, canvas */ img, video, iframe, canvas, svg, picture, [style*="background-image"] { filter: invert(1) hue-rotate(180deg) !important; } /* Preserve specific elements that should not be inverted */ .no-dark-mode, .no-dark-mode *, [data-theme="light"], [data-theme="light"], .ace_editor, .ace_editor *, .CodeMirror, .CodeMirror *, .monaco-editor, .monaco-editor *, .markdown-body pre, .markdown-body pre *, .highlight, .highlight *, pre code, pre code * { filter: none !important; } /* Fix common UI elements */ .modal, .popup, .dropdown-menu, .tooltip, .popover { filter: invert(1) hue-rotate(180deg) !important; background: #2d2d44 !important; border-color: #444 !important; } /* Scrollbars */ ::-webkit-scrollbar { background: #1a1a2e !important; } ::-webkit-scrollbar-thumb { background: #444 !important; } ::-webkit-scrollbar-thumb:hover { background: #555 !important; } /* Selection */ ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; } ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; } '; } function removeDarkMode() { var style = document.getElementById('universal-dark-mode-style'); if (style) style.remove(); } // Toggle with Alt+Shift+D document.addEventListener('keydown', function(e) { if (e.altKey && e.shiftKey && e.key === 'D') { e.preventDefault(); enabled = !enabled; if (enabled) { applyDarkMode(); console.log('[Universal Dark Mode] Enabled'); } else { removeDarkMode(); console.log('[Universal Dark Mode] Disabled'); } } }); // Apply on load applyDarkMode(); // Re-apply on dynamic content var observer = new MutationObserver(function(mutations) { if (enabled && !document.getElementById('universal-dark-mode-style')) { applyDarkMode(); } }); observer.observe(document.head, { childList: true }); console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle'); })(); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })(); Jax bug fixes for the dot product attention by zlsh80826 · Pull Request #236 · NVIDIA/TransformerEngine · GitHub
Skip to content

Jax bug fixes for the dot product attention - #236

Merged
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs
May 23, 2023
Merged

Jax bug fixes for the dot product attention#236
ptrendx merged 8 commits into
NVIDIA:mainfrom
zlsh80826:rewang/fix-no-bias-bugs

Conversation

@zlsh80826

@zlsh80826zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
Collaborator

This PR includes several bug fixes related to the fused and the unfused attention, and also not enabling the fused attention by default.

  1. Fixed an issue where fusing the scale and softmax operations with a bias resulted in incorrect computation flow. The fix ensures that the computation flow becomes Softmax(scale * attn_weights + bias). e594bb6
  2. There is a bug that not clearing S tensor leads to incorrect answer on cuDNN kernel. WAR by clearing S tensor when causal_masking + no_bias and add the relevant unit tests. ac67e91
  3. Enhance the handling of None sharding tensor. 49e73d5
  4. Based on internal discussions, it has been decided to disable fused attention by default for JAX in the upcoming release (0.9). This PR introduces a new flag, NVTE_USE_FUSED_ATTN, which enables fused attention for internal convergence tests (t5x/PAXML). The flag will be removed once the convergence tests are successfully completed. ac67e91
  5. Add thread_local to protect the global static plan cache for thread safety. 8a0d11a

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826zlsh80826 added the bug Something isn't working label May 19, 2023
@mingxu1067

Copy link
Copy Markdown
Collaborator

LGTM

@nouiz

Copy link
Copy Markdown
Collaborator

Can any of those fix affect t5x converges?

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Can any of those fix affect t5x converges?

Yes, the first item is a bug fix for t5x converge, the t5x will not converge (non-fmha, scaledSoftmax) without that fix.
The item 2,3,5 is for PAXML/GPT, not for t5x.

@nouiz

nouiz commented May 19, 2023

Copy link
Copy Markdown
Collaborator

It would be great to have tests for the bugfixes.
I mean, only the second commit have a test.

Comment threadtests/jax/test_fused_attn.py
@nouiz

Copy link
Copy Markdown
Collaborator

Otherwise the small issue, it looks good to me.

Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

zlsh80826 commented May 19, 2023

Copy link
Copy Markdown
CollaboratorAuthor

It would be great to have tests for the bugfixes. I mean, only the second commit have a test.

Yes, I agree. I added an unit test for item 1. (4040ae3). For item 3,5, I will look into whether those can be convered by the multi-card unit tests in the future (other PRs).

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@timmoon10timmoon10 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM

@timmoon10

Copy link
Copy Markdown
Member

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

@jeng1220

jeng1220 commented May 20, 2023

Copy link
Copy Markdown
Contributor

Can you explain why you needed to make the plan caches thread_local? Could you accomplish thread-safety by protecting access with mutexes, or is there a more fundamental reason why threads shouldn't share the same resources? I wonder if the JIT kernel infrastructure needs to make a similar change.

It is according to cuDNN design. The cuDNN handle shouldn't be shared by multiple host threads, and the plans associating with the handle shouldn't either.

We did hit the error if one thread executes a plan generated by another thread before.

https://docs.nvidia.com/deeplearning/cudnn/developer-guide/index.html#thread-safety

@cyanguwacyanguwa left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Looks good to me. Thanks!

q_seqlen = inputs_q.shape[0] if self.transpose_batch_sequence else inputs_q.shape[1]
kv_seqlen = inputs_kv.shape[0] if self.transpose_batch_sequence else inputs_kv.shape[1]
fused_attn_supported_seqlen = [128, 256, 384, 512]
enable_fused_attn = int(os.getenv("NVTE_USE_FUSED_ATTN", "0"))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we use NVTE_FUSED_ATTN instead to keep symmetry with NVTE_FLASH_ATTN?

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Changed in 74eb2b3

Signed-off-by: Reese Wang <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Hello @timmoon10, @cyanguwa,

I think this PR is ready, could you help merge this for 0.9 release? Thanks!

@ptrendx
ptrendx merged commit 6900396 into NVIDIA:mainMay 23, 2023
ptrendx pushed a commit that referenced this pull request May 23, 2023
* Unfused scale+softmax if bias is present
Signed-off-by: Reese Wang <rewang@nvidia.com>
* WAR a causal masking + no_bias bug and add the unittests
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Fix the optional args (bias) sharding
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Disable fused attn in JAX by default, enable it with NVTE_USE_FUSED_ATTN
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add thread local for the plan cache
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Rename dbeta to dbias for the readability
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Add scaled softmax with dropout test cases
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Updated NVTE_FUSED_ATTN variable name
Signed-off-by: Reese Wang <rewang@nvidia.com>
---------
Signed-off-by: Reese Wang <rewang@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

0.9.0bugSomething isn't working

Projects

None yet

Development

Successfully merging this pull request may close these issues.

7 participants

@zlsh80826@mingxu1067@nouiz@timmoon10@jeng1220@cyanguwa@ptrendx