Skip to content

[WIP] Add cudnn fused multi-head attention for JAX - #105

Closed
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha
Closed

[WIP] Add cudnn fused multi-head attention for JAX#105
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha

Conversation

@zlsh80826

Copy link
Copy Markdown
Collaborator
  1. Move scale_factor to core attention to align fused multi-head attention implementation [transformer.py, module.py]
  2. Add cudnn-frontend submodule into 3rdparty/cudnn-frontend
  3. Add cudnn-frontend fused multi-head attention with both self attention and cross attention
  4. Fused multi-head attention will be auto enabled if the network satisify the following rules
    • not decode
    • not transpose_batch_sequence
    • fuse_qkv
    • dropout_rate = 0 (dropout can be fused into FMHA, but we lack model convergence test. Will add it in the future)
    • dtype = bfloat16 or float16
    • q_seqlen and kv_seqlen = 128 or 256 or 384 or 512
  5. MHA supports DP and TP sharding

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

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

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

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

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.

The proposed API is completely different than the rest of TE APIs and is not acceptable.

@cyanguwa is working on a better abstraction for multiple different fused MHA algorithms, please coordinate with her to plug under that API.

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.

OK, I will work with @cyanguwa to see how to integrate the FMHA

@ptrendx
ptrendx requested a review from cyanguwaMarch 16, 2023 17:50
@zlsh80826
zlsh80826 marked this pull request as draft March 20, 2023 12:39
@zlsh80826zlsh80826 changed the title Add cudnn fused multi-head attention for JAX[WIP] Add cudnn fused multi-head attention for JAXMar 20, 2023
jax.config.update('experimental_xmap_spmd_lowering_manual', True)


def self_fmha(qkv: jnp.ndarray,

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.

Is this the public API that people should use? If it is, that is not great:

  • the name is not descriptive - what is "fmha"?
  • what it implements is not "Multihead attention" but rather "Dot Product Attention" from the "Attention is all you need paper", so should not really use mha in the name
  • docstrings are missing
  • it will not extend to the non-zero dropout case since there you need offset too. Even if we do not implement that case right away, we should nto expose the API that will prevent us from doing so in the future.
  • what is "scaling factor"?
  • why isn't it a module like the other high-level APIs?

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.

  • The public API that people should use is MultiHeadAttention(TE API doc), the self_fmha and cross_fmha are custom_calls used internally for JAX-TE instead of public users.
  • The "fmha" naming is inherited by apex, I can rename them to self(cross)_fused_dot_product_attention if they are preferred in TE.
  • As described above, the custom_calls are not intented to public users. The doc strings for MultiHeadAttetion are written on MultiHeadAttention
  • The original cuDNN sample has missed the offset and used a CPU-side seed. But yes, we should use both seed and offset and keep them as the device pointers.
  • Is scaling_factor ambiguous here? From the "Attention Is All You Need paper", the scaling factor is rsqrt(head_dim) that scaled to dot_product(Q, tranpose(K))
  • MultiHeadAttention(TE API doc) is implemented as a module. self_fmha is a custom_call wrapper, it is the same level as fp8_dot

I will add some commits for fmha renaming and seed/offset changes. About the naming, which is preferred in TE? self(cross)_fused_dot_product_attention, self(cross)_fused_attention or other?

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.

Oh boy... This means we have a divergence between pyTorch and JAX APIs. pyTorch doesn't expose the MultiHeadAttention API and now this kind of forces us to do that. @ksivaman for visibility. Still, for consistency I would really like for JAX to expose DotProductAttention to be in line with pyTorch (and then of course the MultiHeadAttention API can use that similarly to how it uses the other modules).

If those are not public APIs then most of my comments are not applicable. I would probably settle for self/cross_fused_attention.

About scaling factor - sorry, my bad, I forgot about this scaling and thought about this qk layer scaling that megatron uses. It's fine as is 👍.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

I would really like for JAX to expose DotProductAttention to be in line with pyTorch

Reese will submit another PR to solve this issue

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.

Yes, I will modularize DotProductAttention for JAX in the coming days before integrating the fused attentions.

@timmoon10

Copy link
Copy Markdown
Member

PyTorch support for fused attention is added in #155.

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Close this PR and I will open a new one for the fp16/bf16 fused attention (max_seq <= 512)

Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants

@zlsh80826@timmoon10@jeng1220@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" + '
[WIP] Add cudnn fused multi-head attention for JAX by zlsh80826 · Pull Request #105 · NVIDIA/TransformerEngine · GitHub
Skip to content

[WIP] Add cudnn fused multi-head attention for JAX - #105

Closed
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha
Closed

[WIP] Add cudnn fused multi-head attention for JAX#105
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha

Conversation

@zlsh80826

Copy link
Copy Markdown
Collaborator
  1. Move scale_factor to core attention to align fused multi-head attention implementation [transformer.py, module.py]
  2. Add cudnn-frontend submodule into 3rdparty/cudnn-frontend
  3. Add cudnn-frontend fused multi-head attention with both self attention and cross attention
  4. Fused multi-head attention will be auto enabled if the network satisify the following rules
    • not decode
    • not transpose_batch_sequence
    • fuse_qkv
    • dropout_rate = 0 (dropout can be fused into FMHA, but we lack model convergence test. Will add it in the future)
    • dtype = bfloat16 or float16
    • q_seqlen and kv_seqlen = 128 or 256 or 384 or 512
  5. MHA supports DP and TP sharding

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

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

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

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

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.

The proposed API is completely different than the rest of TE APIs and is not acceptable.

@cyanguwa is working on a better abstraction for multiple different fused MHA algorithms, please coordinate with her to plug under that API.

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.

OK, I will work with @cyanguwa to see how to integrate the FMHA

@ptrendx
ptrendx requested a review from cyanguwaMarch 16, 2023 17:50
@zlsh80826
zlsh80826 marked this pull request as draft March 20, 2023 12:39
@zlsh80826zlsh80826 changed the title Add cudnn fused multi-head attention for JAX[WIP] Add cudnn fused multi-head attention for JAXMar 20, 2023
jax.config.update('experimental_xmap_spmd_lowering_manual', True)


def self_fmha(qkv: jnp.ndarray,

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.

Is this the public API that people should use? If it is, that is not great:

  • the name is not descriptive - what is "fmha"?
  • what it implements is not "Multihead attention" but rather "Dot Product Attention" from the "Attention is all you need paper", so should not really use mha in the name
  • docstrings are missing
  • it will not extend to the non-zero dropout case since there you need offset too. Even if we do not implement that case right away, we should nto expose the API that will prevent us from doing so in the future.
  • what is "scaling factor"?
  • why isn't it a module like the other high-level APIs?

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.

  • The public API that people should use is MultiHeadAttention(TE API doc), the self_fmha and cross_fmha are custom_calls used internally for JAX-TE instead of public users.
  • The "fmha" naming is inherited by apex, I can rename them to self(cross)_fused_dot_product_attention if they are preferred in TE.
  • As described above, the custom_calls are not intented to public users. The doc strings for MultiHeadAttetion are written on MultiHeadAttention
  • The original cuDNN sample has missed the offset and used a CPU-side seed. But yes, we should use both seed and offset and keep them as the device pointers.
  • Is scaling_factor ambiguous here? From the "Attention Is All You Need paper", the scaling factor is rsqrt(head_dim) that scaled to dot_product(Q, tranpose(K))
  • MultiHeadAttention(TE API doc) is implemented as a module. self_fmha is a custom_call wrapper, it is the same level as fp8_dot

I will add some commits for fmha renaming and seed/offset changes. About the naming, which is preferred in TE? self(cross)_fused_dot_product_attention, self(cross)_fused_attention or other?

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.

Oh boy... This means we have a divergence between pyTorch and JAX APIs. pyTorch doesn't expose the MultiHeadAttention API and now this kind of forces us to do that. @ksivaman for visibility. Still, for consistency I would really like for JAX to expose DotProductAttention to be in line with pyTorch (and then of course the MultiHeadAttention API can use that similarly to how it uses the other modules).

If those are not public APIs then most of my comments are not applicable. I would probably settle for self/cross_fused_attention.

About scaling factor - sorry, my bad, I forgot about this scaling and thought about this qk layer scaling that megatron uses. It's fine as is 👍.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

I would really like for JAX to expose DotProductAttention to be in line with pyTorch

Reese will submit another PR to solve this issue

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.

Yes, I will modularize DotProductAttention for JAX in the coming days before integrating the fused attentions.

@timmoon10

Copy link
Copy Markdown
Member

PyTorch support for fused attention is added in #155.

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Close this PR and I will open a new one for the fp16/bf16 fused attention (max_seq <= 512)

Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants

@zlsh80826@timmoon10@jeng1220@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('^' + ".*" + ' [WIP] Add cudnn fused multi-head attention for JAX by zlsh80826 · Pull Request #105 · NVIDIA/TransformerEngine · GitHub
Skip to content

[WIP] Add cudnn fused multi-head attention for JAX - #105

Closed
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha
Closed

[WIP] Add cudnn fused multi-head attention for JAX#105
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha

Conversation

@zlsh80826

Copy link
Copy Markdown
Collaborator
  1. Move scale_factor to core attention to align fused multi-head attention implementation [transformer.py, module.py]
  2. Add cudnn-frontend submodule into 3rdparty/cudnn-frontend
  3. Add cudnn-frontend fused multi-head attention with both self attention and cross attention
  4. Fused multi-head attention will be auto enabled if the network satisify the following rules
    • not decode
    • not transpose_batch_sequence
    • fuse_qkv
    • dropout_rate = 0 (dropout can be fused into FMHA, but we lack model convergence test. Will add it in the future)
    • dtype = bfloat16 or float16
    • q_seqlen and kv_seqlen = 128 or 256 or 384 or 512
  5. MHA supports DP and TP sharding

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

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

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

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

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.

The proposed API is completely different than the rest of TE APIs and is not acceptable.

@cyanguwa is working on a better abstraction for multiple different fused MHA algorithms, please coordinate with her to plug under that API.

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.

OK, I will work with @cyanguwa to see how to integrate the FMHA

@ptrendx
ptrendx requested a review from cyanguwaMarch 16, 2023 17:50
@zlsh80826
zlsh80826 marked this pull request as draft March 20, 2023 12:39
@zlsh80826zlsh80826 changed the title Add cudnn fused multi-head attention for JAX[WIP] Add cudnn fused multi-head attention for JAXMar 20, 2023
jax.config.update('experimental_xmap_spmd_lowering_manual', True)


def self_fmha(qkv: jnp.ndarray,

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.

Is this the public API that people should use? If it is, that is not great:

  • the name is not descriptive - what is "fmha"?
  • what it implements is not "Multihead attention" but rather "Dot Product Attention" from the "Attention is all you need paper", so should not really use mha in the name
  • docstrings are missing
  • it will not extend to the non-zero dropout case since there you need offset too. Even if we do not implement that case right away, we should nto expose the API that will prevent us from doing so in the future.
  • what is "scaling factor"?
  • why isn't it a module like the other high-level APIs?

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.

  • The public API that people should use is MultiHeadAttention(TE API doc), the self_fmha and cross_fmha are custom_calls used internally for JAX-TE instead of public users.
  • The "fmha" naming is inherited by apex, I can rename them to self(cross)_fused_dot_product_attention if they are preferred in TE.
  • As described above, the custom_calls are not intented to public users. The doc strings for MultiHeadAttetion are written on MultiHeadAttention
  • The original cuDNN sample has missed the offset and used a CPU-side seed. But yes, we should use both seed and offset and keep them as the device pointers.
  • Is scaling_factor ambiguous here? From the "Attention Is All You Need paper", the scaling factor is rsqrt(head_dim) that scaled to dot_product(Q, tranpose(K))
  • MultiHeadAttention(TE API doc) is implemented as a module. self_fmha is a custom_call wrapper, it is the same level as fp8_dot

I will add some commits for fmha renaming and seed/offset changes. About the naming, which is preferred in TE? self(cross)_fused_dot_product_attention, self(cross)_fused_attention or other?

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.

Oh boy... This means we have a divergence between pyTorch and JAX APIs. pyTorch doesn't expose the MultiHeadAttention API and now this kind of forces us to do that. @ksivaman for visibility. Still, for consistency I would really like for JAX to expose DotProductAttention to be in line with pyTorch (and then of course the MultiHeadAttention API can use that similarly to how it uses the other modules).

If those are not public APIs then most of my comments are not applicable. I would probably settle for self/cross_fused_attention.

About scaling factor - sorry, my bad, I forgot about this scaling and thought about this qk layer scaling that megatron uses. It's fine as is 👍.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

I would really like for JAX to expose DotProductAttention to be in line with pyTorch

Reese will submit another PR to solve this issue

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.

Yes, I will modularize DotProductAttention for JAX in the coming days before integrating the fused attentions.

@timmoon10

Copy link
Copy Markdown
Member

PyTorch support for fused attention is added in #155.

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Close this PR and I will open a new one for the fp16/bf16 fused attention (max_seq <= 512)

Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants

@zlsh80826@timmoon10@jeng1220@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('^' + ".*" + ' [WIP] Add cudnn fused multi-head attention for JAX by zlsh80826 · Pull Request #105 · NVIDIA/TransformerEngine · GitHub
Skip to content

[WIP] Add cudnn fused multi-head attention for JAX - #105

Closed
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha
Closed

[WIP] Add cudnn fused multi-head attention for JAX#105
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha

Conversation

@zlsh80826

Copy link
Copy Markdown
Collaborator
  1. Move scale_factor to core attention to align fused multi-head attention implementation [transformer.py, module.py]
  2. Add cudnn-frontend submodule into 3rdparty/cudnn-frontend
  3. Add cudnn-frontend fused multi-head attention with both self attention and cross attention
  4. Fused multi-head attention will be auto enabled if the network satisify the following rules
    • not decode
    • not transpose_batch_sequence
    • fuse_qkv
    • dropout_rate = 0 (dropout can be fused into FMHA, but we lack model convergence test. Will add it in the future)
    • dtype = bfloat16 or float16
    • q_seqlen and kv_seqlen = 128 or 256 or 384 or 512
  5. MHA supports DP and TP sharding

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

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

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

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

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.

The proposed API is completely different than the rest of TE APIs and is not acceptable.

@cyanguwa is working on a better abstraction for multiple different fused MHA algorithms, please coordinate with her to plug under that API.

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.

OK, I will work with @cyanguwa to see how to integrate the FMHA

@ptrendx
ptrendx requested a review from cyanguwaMarch 16, 2023 17:50
@zlsh80826
zlsh80826 marked this pull request as draft March 20, 2023 12:39
@zlsh80826zlsh80826 changed the title Add cudnn fused multi-head attention for JAX[WIP] Add cudnn fused multi-head attention for JAXMar 20, 2023
jax.config.update('experimental_xmap_spmd_lowering_manual', True)


def self_fmha(qkv: jnp.ndarray,

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.

Is this the public API that people should use? If it is, that is not great:

  • the name is not descriptive - what is "fmha"?
  • what it implements is not "Multihead attention" but rather "Dot Product Attention" from the "Attention is all you need paper", so should not really use mha in the name
  • docstrings are missing
  • it will not extend to the non-zero dropout case since there you need offset too. Even if we do not implement that case right away, we should nto expose the API that will prevent us from doing so in the future.
  • what is "scaling factor"?
  • why isn't it a module like the other high-level APIs?

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.

  • The public API that people should use is MultiHeadAttention(TE API doc), the self_fmha and cross_fmha are custom_calls used internally for JAX-TE instead of public users.
  • The "fmha" naming is inherited by apex, I can rename them to self(cross)_fused_dot_product_attention if they are preferred in TE.
  • As described above, the custom_calls are not intented to public users. The doc strings for MultiHeadAttetion are written on MultiHeadAttention
  • The original cuDNN sample has missed the offset and used a CPU-side seed. But yes, we should use both seed and offset and keep them as the device pointers.
  • Is scaling_factor ambiguous here? From the "Attention Is All You Need paper", the scaling factor is rsqrt(head_dim) that scaled to dot_product(Q, tranpose(K))
  • MultiHeadAttention(TE API doc) is implemented as a module. self_fmha is a custom_call wrapper, it is the same level as fp8_dot

I will add some commits for fmha renaming and seed/offset changes. About the naming, which is preferred in TE? self(cross)_fused_dot_product_attention, self(cross)_fused_attention or other?

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.

Oh boy... This means we have a divergence between pyTorch and JAX APIs. pyTorch doesn't expose the MultiHeadAttention API and now this kind of forces us to do that. @ksivaman for visibility. Still, for consistency I would really like for JAX to expose DotProductAttention to be in line with pyTorch (and then of course the MultiHeadAttention API can use that similarly to how it uses the other modules).

If those are not public APIs then most of my comments are not applicable. I would probably settle for self/cross_fused_attention.

About scaling factor - sorry, my bad, I forgot about this scaling and thought about this qk layer scaling that megatron uses. It's fine as is 👍.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

I would really like for JAX to expose DotProductAttention to be in line with pyTorch

Reese will submit another PR to solve this issue

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.

Yes, I will modularize DotProductAttention for JAX in the coming days before integrating the fused attentions.

@timmoon10

Copy link
Copy Markdown
Member

PyTorch support for fused attention is added in #155.

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Close this PR and I will open a new one for the fp16/bf16 fused attention (max_seq <= 512)

Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants

@zlsh80826@timmoon10@jeng1220@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" + ' [WIP] Add cudnn fused multi-head attention for JAX by zlsh80826 · Pull Request #105 · NVIDIA/TransformerEngine · GitHub
Skip to content

[WIP] Add cudnn fused multi-head attention for JAX - #105

Closed
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha
Closed

[WIP] Add cudnn fused multi-head attention for JAX#105
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha

Conversation

@zlsh80826

Copy link
Copy Markdown
Collaborator
  1. Move scale_factor to core attention to align fused multi-head attention implementation [transformer.py, module.py]
  2. Add cudnn-frontend submodule into 3rdparty/cudnn-frontend
  3. Add cudnn-frontend fused multi-head attention with both self attention and cross attention
  4. Fused multi-head attention will be auto enabled if the network satisify the following rules
    • not decode
    • not transpose_batch_sequence
    • fuse_qkv
    • dropout_rate = 0 (dropout can be fused into FMHA, but we lack model convergence test. Will add it in the future)
    • dtype = bfloat16 or float16
    • q_seqlen and kv_seqlen = 128 or 256 or 384 or 512
  5. MHA supports DP and TP sharding

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

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

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

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

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.

The proposed API is completely different than the rest of TE APIs and is not acceptable.

@cyanguwa is working on a better abstraction for multiple different fused MHA algorithms, please coordinate with her to plug under that API.

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.

OK, I will work with @cyanguwa to see how to integrate the FMHA

@ptrendx
ptrendx requested a review from cyanguwaMarch 16, 2023 17:50
@zlsh80826
zlsh80826 marked this pull request as draft March 20, 2023 12:39
@zlsh80826zlsh80826 changed the title Add cudnn fused multi-head attention for JAX[WIP] Add cudnn fused multi-head attention for JAXMar 20, 2023
jax.config.update('experimental_xmap_spmd_lowering_manual', True)


def self_fmha(qkv: jnp.ndarray,

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.

Is this the public API that people should use? If it is, that is not great:

  • the name is not descriptive - what is "fmha"?
  • what it implements is not "Multihead attention" but rather "Dot Product Attention" from the "Attention is all you need paper", so should not really use mha in the name
  • docstrings are missing
  • it will not extend to the non-zero dropout case since there you need offset too. Even if we do not implement that case right away, we should nto expose the API that will prevent us from doing so in the future.
  • what is "scaling factor"?
  • why isn't it a module like the other high-level APIs?

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.

  • The public API that people should use is MultiHeadAttention(TE API doc), the self_fmha and cross_fmha are custom_calls used internally for JAX-TE instead of public users.
  • The "fmha" naming is inherited by apex, I can rename them to self(cross)_fused_dot_product_attention if they are preferred in TE.
  • As described above, the custom_calls are not intented to public users. The doc strings for MultiHeadAttetion are written on MultiHeadAttention
  • The original cuDNN sample has missed the offset and used a CPU-side seed. But yes, we should use both seed and offset and keep them as the device pointers.
  • Is scaling_factor ambiguous here? From the "Attention Is All You Need paper", the scaling factor is rsqrt(head_dim) that scaled to dot_product(Q, tranpose(K))
  • MultiHeadAttention(TE API doc) is implemented as a module. self_fmha is a custom_call wrapper, it is the same level as fp8_dot

I will add some commits for fmha renaming and seed/offset changes. About the naming, which is preferred in TE? self(cross)_fused_dot_product_attention, self(cross)_fused_attention or other?

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.

Oh boy... This means we have a divergence between pyTorch and JAX APIs. pyTorch doesn't expose the MultiHeadAttention API and now this kind of forces us to do that. @ksivaman for visibility. Still, for consistency I would really like for JAX to expose DotProductAttention to be in line with pyTorch (and then of course the MultiHeadAttention API can use that similarly to how it uses the other modules).

If those are not public APIs then most of my comments are not applicable. I would probably settle for self/cross_fused_attention.

About scaling factor - sorry, my bad, I forgot about this scaling and thought about this qk layer scaling that megatron uses. It's fine as is 👍.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

I would really like for JAX to expose DotProductAttention to be in line with pyTorch

Reese will submit another PR to solve this issue

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.

Yes, I will modularize DotProductAttention for JAX in the coming days before integrating the fused attentions.

@timmoon10

Copy link
Copy Markdown
Member

PyTorch support for fused attention is added in #155.

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Close this PR and I will open a new one for the fp16/bf16 fused attention (max_seq <= 512)

Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants

@zlsh80826@timmoon10@jeng1220@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('^' + ".*" + ' [WIP] Add cudnn fused multi-head attention for JAX by zlsh80826 · Pull Request #105 · NVIDIA/TransformerEngine · GitHub
Skip to content

[WIP] Add cudnn fused multi-head attention for JAX - #105

Closed
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha
Closed

[WIP] Add cudnn fused multi-head attention for JAX#105
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha

Conversation

@zlsh80826

Copy link
Copy Markdown
Collaborator
  1. Move scale_factor to core attention to align fused multi-head attention implementation [transformer.py, module.py]
  2. Add cudnn-frontend submodule into 3rdparty/cudnn-frontend
  3. Add cudnn-frontend fused multi-head attention with both self attention and cross attention
  4. Fused multi-head attention will be auto enabled if the network satisify the following rules
    • not decode
    • not transpose_batch_sequence
    • fuse_qkv
    • dropout_rate = 0 (dropout can be fused into FMHA, but we lack model convergence test. Will add it in the future)
    • dtype = bfloat16 or float16
    • q_seqlen and kv_seqlen = 128 or 256 or 384 or 512
  5. MHA supports DP and TP sharding

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

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

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

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

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.

The proposed API is completely different than the rest of TE APIs and is not acceptable.

@cyanguwa is working on a better abstraction for multiple different fused MHA algorithms, please coordinate with her to plug under that API.

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.

OK, I will work with @cyanguwa to see how to integrate the FMHA

@ptrendx
ptrendx requested a review from cyanguwaMarch 16, 2023 17:50
@zlsh80826
zlsh80826 marked this pull request as draft March 20, 2023 12:39
@zlsh80826zlsh80826 changed the title Add cudnn fused multi-head attention for JAX[WIP] Add cudnn fused multi-head attention for JAXMar 20, 2023
jax.config.update('experimental_xmap_spmd_lowering_manual', True)


def self_fmha(qkv: jnp.ndarray,

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.

Is this the public API that people should use? If it is, that is not great:

  • the name is not descriptive - what is "fmha"?
  • what it implements is not "Multihead attention" but rather "Dot Product Attention" from the "Attention is all you need paper", so should not really use mha in the name
  • docstrings are missing
  • it will not extend to the non-zero dropout case since there you need offset too. Even if we do not implement that case right away, we should nto expose the API that will prevent us from doing so in the future.
  • what is "scaling factor"?
  • why isn't it a module like the other high-level APIs?

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.

  • The public API that people should use is MultiHeadAttention(TE API doc), the self_fmha and cross_fmha are custom_calls used internally for JAX-TE instead of public users.
  • The "fmha" naming is inherited by apex, I can rename them to self(cross)_fused_dot_product_attention if they are preferred in TE.
  • As described above, the custom_calls are not intented to public users. The doc strings for MultiHeadAttetion are written on MultiHeadAttention
  • The original cuDNN sample has missed the offset and used a CPU-side seed. But yes, we should use both seed and offset and keep them as the device pointers.
  • Is scaling_factor ambiguous here? From the "Attention Is All You Need paper", the scaling factor is rsqrt(head_dim) that scaled to dot_product(Q, tranpose(K))
  • MultiHeadAttention(TE API doc) is implemented as a module. self_fmha is a custom_call wrapper, it is the same level as fp8_dot

I will add some commits for fmha renaming and seed/offset changes. About the naming, which is preferred in TE? self(cross)_fused_dot_product_attention, self(cross)_fused_attention or other?

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.

Oh boy... This means we have a divergence between pyTorch and JAX APIs. pyTorch doesn't expose the MultiHeadAttention API and now this kind of forces us to do that. @ksivaman for visibility. Still, for consistency I would really like for JAX to expose DotProductAttention to be in line with pyTorch (and then of course the MultiHeadAttention API can use that similarly to how it uses the other modules).

If those are not public APIs then most of my comments are not applicable. I would probably settle for self/cross_fused_attention.

About scaling factor - sorry, my bad, I forgot about this scaling and thought about this qk layer scaling that megatron uses. It's fine as is 👍.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

I would really like for JAX to expose DotProductAttention to be in line with pyTorch

Reese will submit another PR to solve this issue

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.

Yes, I will modularize DotProductAttention for JAX in the coming days before integrating the fused attentions.

@timmoon10

Copy link
Copy Markdown
Member

PyTorch support for fused attention is added in #155.

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Close this PR and I will open a new one for the fp16/bf16 fused attention (max_seq <= 512)

Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants

@zlsh80826@timmoon10@jeng1220@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('^' + ".*" + ' [WIP] Add cudnn fused multi-head attention for JAX by zlsh80826 · Pull Request #105 · NVIDIA/TransformerEngine · GitHub
Skip to content

[WIP] Add cudnn fused multi-head attention for JAX - #105

Closed
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha
Closed

[WIP] Add cudnn fused multi-head attention for JAX#105
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha

Conversation

@zlsh80826

Copy link
Copy Markdown
Collaborator
  1. Move scale_factor to core attention to align fused multi-head attention implementation [transformer.py, module.py]
  2. Add cudnn-frontend submodule into 3rdparty/cudnn-frontend
  3. Add cudnn-frontend fused multi-head attention with both self attention and cross attention
  4. Fused multi-head attention will be auto enabled if the network satisify the following rules
    • not decode
    • not transpose_batch_sequence
    • fuse_qkv
    • dropout_rate = 0 (dropout can be fused into FMHA, but we lack model convergence test. Will add it in the future)
    • dtype = bfloat16 or float16
    • q_seqlen and kv_seqlen = 128 or 256 or 384 or 512
  5. MHA supports DP and TP sharding

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

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

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

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

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.

The proposed API is completely different than the rest of TE APIs and is not acceptable.

@cyanguwa is working on a better abstraction for multiple different fused MHA algorithms, please coordinate with her to plug under that API.

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.

OK, I will work with @cyanguwa to see how to integrate the FMHA

@ptrendx
ptrendx requested a review from cyanguwaMarch 16, 2023 17:50
@zlsh80826
zlsh80826 marked this pull request as draft March 20, 2023 12:39
@zlsh80826zlsh80826 changed the title Add cudnn fused multi-head attention for JAX[WIP] Add cudnn fused multi-head attention for JAXMar 20, 2023
jax.config.update('experimental_xmap_spmd_lowering_manual', True)


def self_fmha(qkv: jnp.ndarray,

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.

Is this the public API that people should use? If it is, that is not great:

  • the name is not descriptive - what is "fmha"?
  • what it implements is not "Multihead attention" but rather "Dot Product Attention" from the "Attention is all you need paper", so should not really use mha in the name
  • docstrings are missing
  • it will not extend to the non-zero dropout case since there you need offset too. Even if we do not implement that case right away, we should nto expose the API that will prevent us from doing so in the future.
  • what is "scaling factor"?
  • why isn't it a module like the other high-level APIs?

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.

  • The public API that people should use is MultiHeadAttention(TE API doc), the self_fmha and cross_fmha are custom_calls used internally for JAX-TE instead of public users.
  • The "fmha" naming is inherited by apex, I can rename them to self(cross)_fused_dot_product_attention if they are preferred in TE.
  • As described above, the custom_calls are not intented to public users. The doc strings for MultiHeadAttetion are written on MultiHeadAttention
  • The original cuDNN sample has missed the offset and used a CPU-side seed. But yes, we should use both seed and offset and keep them as the device pointers.
  • Is scaling_factor ambiguous here? From the "Attention Is All You Need paper", the scaling factor is rsqrt(head_dim) that scaled to dot_product(Q, tranpose(K))
  • MultiHeadAttention(TE API doc) is implemented as a module. self_fmha is a custom_call wrapper, it is the same level as fp8_dot

I will add some commits for fmha renaming and seed/offset changes. About the naming, which is preferred in TE? self(cross)_fused_dot_product_attention, self(cross)_fused_attention or other?

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.

Oh boy... This means we have a divergence between pyTorch and JAX APIs. pyTorch doesn't expose the MultiHeadAttention API and now this kind of forces us to do that. @ksivaman for visibility. Still, for consistency I would really like for JAX to expose DotProductAttention to be in line with pyTorch (and then of course the MultiHeadAttention API can use that similarly to how it uses the other modules).

If those are not public APIs then most of my comments are not applicable. I would probably settle for self/cross_fused_attention.

About scaling factor - sorry, my bad, I forgot about this scaling and thought about this qk layer scaling that megatron uses. It's fine as is 👍.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

I would really like for JAX to expose DotProductAttention to be in line with pyTorch

Reese will submit another PR to solve this issue

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.

Yes, I will modularize DotProductAttention for JAX in the coming days before integrating the fused attentions.

@timmoon10

Copy link
Copy Markdown
Member

PyTorch support for fused attention is added in #155.

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Close this PR and I will open a new one for the fp16/bf16 fused attention (max_seq <= 512)

Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants

@zlsh80826@timmoon10@jeng1220@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); } })(); })(); [WIP] Add cudnn fused multi-head attention for JAX by zlsh80826 · Pull Request #105 · NVIDIA/TransformerEngine · GitHub
Skip to content

[WIP] Add cudnn fused multi-head attention for JAX - #105

Closed
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha
Closed

[WIP] Add cudnn fused multi-head attention for JAX#105
zlsh80826 wants to merge 7 commits into
NVIDIA:mainfrom
zlsh80826:rewang/add-fmha

Conversation

@zlsh80826

Copy link
Copy Markdown
Collaborator
  1. Move scale_factor to core attention to align fused multi-head attention implementation [transformer.py, module.py]
  2. Add cudnn-frontend submodule into 3rdparty/cudnn-frontend
  3. Add cudnn-frontend fused multi-head attention with both self attention and cross attention
  4. Fused multi-head attention will be auto enabled if the network satisify the following rules
    • not decode
    • not transpose_batch_sequence
    • fuse_qkv
    • dropout_rate = 0 (dropout can be fused into FMHA, but we lack model convergence test. Will add it in the future)
    • dtype = bfloat16 or float16
    • q_seqlen and kv_seqlen = 128 or 256 or 384 or 512
  5. MHA supports DP and TP sharding

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

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

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

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

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.

The proposed API is completely different than the rest of TE APIs and is not acceptable.

@cyanguwa is working on a better abstraction for multiple different fused MHA algorithms, please coordinate with her to plug under that API.

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.

OK, I will work with @cyanguwa to see how to integrate the FMHA

@ptrendx
ptrendx requested a review from cyanguwaMarch 16, 2023 17:50
@zlsh80826
zlsh80826 marked this pull request as draft March 20, 2023 12:39
@zlsh80826zlsh80826 changed the title Add cudnn fused multi-head attention for JAX[WIP] Add cudnn fused multi-head attention for JAXMar 20, 2023
jax.config.update('experimental_xmap_spmd_lowering_manual', True)


def self_fmha(qkv: jnp.ndarray,

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.

Is this the public API that people should use? If it is, that is not great:

  • the name is not descriptive - what is "fmha"?
  • what it implements is not "Multihead attention" but rather "Dot Product Attention" from the "Attention is all you need paper", so should not really use mha in the name
  • docstrings are missing
  • it will not extend to the non-zero dropout case since there you need offset too. Even if we do not implement that case right away, we should nto expose the API that will prevent us from doing so in the future.
  • what is "scaling factor"?
  • why isn't it a module like the other high-level APIs?

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.

  • The public API that people should use is MultiHeadAttention(TE API doc), the self_fmha and cross_fmha are custom_calls used internally for JAX-TE instead of public users.
  • The "fmha" naming is inherited by apex, I can rename them to self(cross)_fused_dot_product_attention if they are preferred in TE.
  • As described above, the custom_calls are not intented to public users. The doc strings for MultiHeadAttetion are written on MultiHeadAttention
  • The original cuDNN sample has missed the offset and used a CPU-side seed. But yes, we should use both seed and offset and keep them as the device pointers.
  • Is scaling_factor ambiguous here? From the "Attention Is All You Need paper", the scaling factor is rsqrt(head_dim) that scaled to dot_product(Q, tranpose(K))
  • MultiHeadAttention(TE API doc) is implemented as a module. self_fmha is a custom_call wrapper, it is the same level as fp8_dot

I will add some commits for fmha renaming and seed/offset changes. About the naming, which is preferred in TE? self(cross)_fused_dot_product_attention, self(cross)_fused_attention or other?

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.

Oh boy... This means we have a divergence between pyTorch and JAX APIs. pyTorch doesn't expose the MultiHeadAttention API and now this kind of forces us to do that. @ksivaman for visibility. Still, for consistency I would really like for JAX to expose DotProductAttention to be in line with pyTorch (and then of course the MultiHeadAttention API can use that similarly to how it uses the other modules).

If those are not public APIs then most of my comments are not applicable. I would probably settle for self/cross_fused_attention.

About scaling factor - sorry, my bad, I forgot about this scaling and thought about this qk layer scaling that megatron uses. It's fine as is 👍.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

I would really like for JAX to expose DotProductAttention to be in line with pyTorch

Reese will submit another PR to solve this issue

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.

Yes, I will modularize DotProductAttention for JAX in the coming days before integrating the fused attentions.

@timmoon10

Copy link
Copy Markdown
Member

PyTorch support for fused attention is added in #155.

@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

Close this PR and I will open a new one for the fp16/bf16 fused attention (max_seq <= 512)

Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants

@zlsh80826@timmoon10@jeng1220@ptrendx