') + ')', '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('^' + ".*" + ', '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" + ', '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('^' + ".*" + ', '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); } })(); })(); Expand FA coverage to padding masks by ksivaman · Pull Request #291 · NVIDIA/TransformerEngine · GitHub
Skip to content

Expand FA coverage to padding masks - #291

Closed
ksivaman wants to merge 12 commits into
NVIDIA:mainfrom
ksivaman:expand_flash_attn_coverage
Closed

Expand FA coverage to padding masks#291
ksivaman wants to merge 12 commits into
NVIDIA:mainfrom
ksivaman:expand_flash_attn_coverage

Conversation

@ksivaman

Copy link
Copy Markdown
Member

No description provided.

Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
@ksivamanksivaman self-assigned this Jun 21, 2023
@ksivaman
ksivaman marked this pull request as draft June 21, 2023 19:42
@ksivaman
ksivaman marked this pull request as ready for review June 21, 2023 20:03
Comment threadtransformer_engine/pytorch/softmax.py Outdated
Comment threadtransformer_engine/pytorch/softmax.py Outdated
Comment on lines +264 to +269
if dtype == torch.float32:
mask_shape = torch.Size([1, 1, config.seq_len, config.seq_len])
else:
mask_shape = torch.Size([config.seq_len, bs])

te_inp_attn_mask = torch.rand(mask_shape).cuda().bool()

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.

Why does the FP32 test require a 4D mask while other dtypes use a 2D mask?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

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

The FP32 doesn't take the FA path and so uses the PyTorch torch softmax path for padding mask which expects this

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.

Ah, so FA vs PyT is no longer an implementation detail since it expects a different mask format. This makes me think we need an option to explicitly enable or disable FA, and to error out instead of falling back to PyT.

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.

Alternatively, we should expect the mask to be in the format for PyT, and then convert it to the FA format internally so it isn't visible to the user.

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.

Agree with Tim. OR, accept both in both cases and generate the right one for the current implementation.

To be honest though, the current API for mask type is just bad. For example, what if you want padding AND causal? Then your only option is set the padding mask type and create the causal mask yourself (which would mean you can't do FA in the current implementation which is bad). There is also nothing stopping you really from having arbitrary padding mask (as in, with random elements zeroed out), which would break the assumptions you have in this PR. I think what we could do is introduce new names for the padding type (let's say "pad", "arbitrary" and "no_mask"), deprecate the old names and add causal as separate switch. Then we could say that pad only accepts the "mask" that is actually just list of sequence lengths, arbitrary is whatever you want (and will not go through FA) and causal could be switched irrespective of the mask type.

Comment threadqa/L0_lint/pylintrc Outdated
@timmoon10
timmoon10 self-requested a review June 21, 2023 21:20
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
@ksivaman

Copy link
Copy Markdown
MemberAuthor

/te-ci

Comment on lines +264 to +269
if dtype == torch.float32:
mask_shape = torch.Size([1, 1, config.seq_len, config.seq_len])
else:
mask_shape = torch.Size([config.seq_len, bs])

te_inp_attn_mask = torch.rand(mask_shape).cuda().bool()

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.

Ah, so FA vs PyT is no longer an implementation detail since it expects a different mask format. This makes me think we need an option to explicitly enable or disable FA, and to error out instead of falling back to PyT.

Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
@timmoon10
timmoon10 self-requested a review June 22, 2023 20:39
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
@jit_fuser
def get_cu_seqlens(padding_mask: torch.Tensor) -> torch.Tensor:
"""
Given a padding mask of shape [seq_len, batch_size], returns an int32

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.

I guess by seq_len you mean max_seqlen? Also, should the mask tensor be required to be a CUDA tensor?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

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

Yeah it is the max_seqlen.

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

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

Yeah it should be cuda tensor since the inp is cuda as well.

if self.attn_mask_type == "padding":
assert (
attention_mask is not None
), "Boolean attention mask must be provided for padding."

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.

So we're settled on requiring a mask tensor than a cu_seqlen tensor from users?

@ksivamanksivamanJun 22, 2023

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

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

Yeah I think so, just to be consistent


use_flash_attention = self.use_flash_attention
if (query_layer.dtype not in [torch.bfloat16, torch.float16]
if (query_layer.dtype not in [torch.bfloat16, torch.float16] # pylint: disable=too-many-boolean-expressions

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.

Could probably use any(x.dtype not in [torch.bfloat16, torch.float16] for x in [query_layer, key_layer, value_layer]). I think pylint recognizes this as one boolean expression.

key_layer,
value_layer)
return self.flash_attention(query_layer, key_layer, value_layer)
attention_func = self.flash_attention if use_flash_attention else self.unfused_attention

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.

Should we include self.fused_attention in this logic as well? In my PR, I kept the selection order as flash -> fused -> unfused.

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

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

Yes I like this hierarchy. I'm switching to it as I resolve all the conflicts

if self.attn_mask_type == "padding":
assert (
attention_mask is not None
), "Boolean attention mask must be provided for padding."

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.

The m512 version of the fused attention supports padding as well, but it doesn't require a mask tensor. Unfortunately which backend to use for fused attention is only determined later in DPA forward. So maybe we can move this check to the FlashAttention module instead.

@ksivamanksivamanJun 22, 2023

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

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

How does it unpad without a mask tensor provided?

Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
@ksivaman
ksivaman marked this pull request as draft June 26, 2023 15:06
@ksivaman

Copy link
Copy Markdown
MemberAuthor

Closing in favor of #302

@ksivaman
ksivaman deleted the expand_flash_attn_coverage branch July 19, 2023 01:41
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

@ksivaman@timmoon10@ptrendx@cyanguwa