') + ')', '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); } })(); })(); [JAX] Fully remove attn_type and set self_attn_mask_type default to 'causal' by zlsh80826 · Pull Request #324 · NVIDIA/TransformerEngine · GitHub
Skip to content

[JAX] Fully remove attn_type and set self_attn_mask_type default to 'causal' - #324

Merged
ksivaman merged 6 commits into
NVIDIA:mainfrom
zlsh80826:rewang/change-self-attn-mask-type-default
Jul 18, 2023
Merged

[JAX] Fully remove attn_type and set self_attn_mask_type default to 'causal'#324
ksivaman merged 6 commits into
NVIDIA:mainfrom
zlsh80826:rewang/change-self-attn-mask-type-default

Conversation

@zlsh80826

Copy link
Copy Markdown
Collaborator

Background:
PR-273:Add self_attn_mask_type and replace attn_type introduced self_attn_mask_type and planed to fully remove attn_type in v0.11.

This PR does the following thing as PR-273 follow up:

  1. Fully remove attn_type
  2. Set the default value of self_attn_mask_type to "causal" to align pytorch framework

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

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

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

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@zlsh80826zlsh80826 changed the title Fully remove attn_type and set self_attn_mask_type default to 'causal'[JAX] Fully remove attn_type and set self_attn_mask_type default to 'causal'Jul 17, 2023
@zlsh80826zlsh80826 changed the title [JAX] Fully remove attn_type and set self_attn_mask_type default to 'causal'[JAX][v0.11] Fully remove attn_type and set self_attn_mask_type default to 'causal'Jul 17, 2023
@ksivamanksivaman changed the title [JAX][v0.11] Fully remove attn_type and set self_attn_mask_type default to 'causal'[JAX] Fully remove attn_type and set self_attn_mask_type default to 'causal'Jul 17, 2023
Comment threadtransformer_engine/jax/flax/transformer.py Outdated
Comment threadtransformer_engine/jax/flax/transformer.py Outdated

@ksivamanksivaman 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 aside from minor doc suggestions

@ksivaman

Copy link
Copy Markdown
Member

/te-ci

zlsh80826and others added 2 commits July 18, 2023 11:21
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: zlsh80826 <rewang@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: zlsh80826 <rewang@nvidia.com>
@zlsh80826

Copy link
Copy Markdown
CollaboratorAuthor

/te-ci

@ksivaman
ksivaman merged commit a3e4e61 into NVIDIA:mainJul 18, 2023
ksivaman added a commit that referenced this pull request Jul 31, 2023
…causal' (#324)
* Fully remove attn_type and set self_attn_mask_type default to 'causal'
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Fix tests with new arguments
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Explicit self_attn_mask_type for examples
Signed-off-by: Reese Wang <rewang@nvidia.com>
* Update transformer_engine/jax/flax/transformer.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: zlsh80826 <rewang@nvidia.com>
* Update transformer_engine/jax/flax/transformer.py
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: zlsh80826 <rewang@nvidia.com>
---------
Signed-off-by: Reese Wang <rewang@nvidia.com>
Signed-off-by: zlsh80826 <rewang@nvidia.com>
Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants

@zlsh80826@ksivaman@timmoon10