Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 75 additions & 0 deletions examples/jax/encoder/test_single_gpu_bf16_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,75 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with BF16 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)

encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
optimizer = optax.sgd(0.001, 0.9)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
99 changes: 99 additions & 0 deletions examples/jax/encoder/test_single_gpu_fp8_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,99 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with FP8 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from cuda import cudart
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te
from transformer_engine.jax.fp8 import FP8Helper
from transformer_engine.common.recipe import Format as FP8Format
from transformer_engine.common.recipe import DelayedScaling

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def gpu_has_fp8():
"""GPU arch has to support FP8"""
cudaSuccess = cudart.cudaError_t.cudaSuccess
ret, gpu_id = cudart.cudaGetDevice()
assert ret == cudaSuccess
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMajor
_, major = cudart.cudaDeviceGetAttribute(flag, gpu_id)
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMinor
_, minor = cudart.cudaDeviceGetAttribute(flag, gpu_id)
sm_arch = major * 10 + minor
return sm_arch >= 89


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
others = FP8Helper.update_fp8_metas(grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
if gpu_has_fp8() is False:
print("GPU doesn't support FP8")
return

rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)
optimizer = optax.sgd(0.001, 0.9)

with te.fp8_autocast(enabled=True, fp8_recipe=DelayedScaling(fp8_format=FP8Format.HYBRID)):
encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)
Comment thread
timmoon10 marked this conversation as resolved.
assert "fp8" in str(jax.make_jaxpr(jitted_train_step)(inputs, state, variables))

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
1 change: 1 addition & 0 deletions qa/L0_jax_unittest/test.sh
Original file line numberDiff line numberDiff line change
Expand Up@@ -6,3 +6,4 @@ set -xe

: ${TE_PATH:=/opt/transformerengine}
pytest -Wignore -v $TE_PATH/tests/jax
pytest -Wignore -v $TE_PATH/examples/jax
38 changes: 17 additions & 21 deletions tests/jax/test_custom_call_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -60,15 +60,15 @@ def test_compile_bf16(self):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
# x = input, matrix 2d
# y = input, matrix 2d (weight)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -84,13 +84,13 @@ def test_compile_fp8(self, compute_type):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -104,13 +104,13 @@ def test_forward_bf16(self, m, n, k):
b = jax.random.normal(subkeys[1], (k, n), jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None))
primitive_out = fp8_dot(fp8_gemm_pkg, *_format2dtypes(None))
ref_out = jnp.dot(a, b)

assert_allclose(primitive_out, ref_out)
Expand All@@ -128,20 +128,20 @@ def test_forward_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

# calculate amax
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
# calculate scale by amax
fp8_meta = FP8Helper._update_fp8_metas_impl(fp8_meta)

fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
ref_out = jnp.dot(a, b)

ref_out = ref_out.astype(jnp.float32)
Expand All@@ -158,13 +158,13 @@ def test_grad_bf16(self, m, n, k):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.mean(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.mean(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

def ref_func(x, y):
return jnp.mean(jnp.dot(x, y))
Expand DownExpand Up@@ -193,15 +193,15 @@ def test_grad_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

def primitive_func(x, y, metas):
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], *metas)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

def ref_func(x, y):
return jnp.sum(jnp.dot(x, y))
Expand DownExpand Up@@ -232,13 +232,13 @@ def test_contracting_dims_bf16(self):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None), ((2, 3), (0, 1))))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None), ((2, 3), (0, 1))))

def ref_func(x, y):
return jnp.sum(lax.dot_general(x, y, dimension_numbers=(((2, 3), (0, 1)), ((), ()))))
Expand DownExpand Up@@ -266,7 +266,7 @@ def test_grad_fp8_mlp_randint(self, m, n, k):
s = jax.random.uniform(subkeys[3], (k,), jnp.bfloat16, 5, 8)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM * 2)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
Expand All@@ -283,7 +283,6 @@ def primitive_func(x, ln_s, y, z, metas):
ln_s,
None,
"rmsnorm",
0,
*compute_type,
activations=activations))

Expand All@@ -305,7 +304,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
amax: jnp.ndarray,
scale: jnp.ndarray,
scale_inv: jnp.ndarray,
amax_history_idx: int,
fwd_dtype,
bwd_dtype,
epsilon=1e-6,
Expand All@@ -323,7 +321,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[:FP8Helper.NUM_META_PER_GEMM],
scale_inv[:FP8Helper.NUM_META_PER_GEMM])
linear_1_out = fp8_dot(fp8_gemm_1_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -341,7 +338,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[FP8Helper.NUM_META_PER_GEMM:],
scale_inv[FP8Helper.NUM_META_PER_GEMM:])
output = fp8_dot(fp8_gemm_2_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -350,7 +346,7 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,

def ref_func(x, ln_s, y, z, metas):
return jnp.mean(
fp8_ln_mlp_py(x, ln_s, y, z, *metas, 0, *compute_type, activations=activations))
fp8_ln_mlp_py(x, ln_s, y, z, *metas, *compute_type, activations=activations))

value_n_grad_primitive_func = jit(value_and_grad(primitive_func, (0, 1, 2, 3)))
value_n_grad_ref_func = jit(value_and_grad(ref_func, (0, 1, 2, 3)))
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all
 blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n 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;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks");
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 75 additions & 0 deletions examples/jax/encoder/test_single_gpu_bf16_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,75 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with BF16 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)

encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
optimizer = optax.sgd(0.001, 0.9)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
99 changes: 99 additions & 0 deletions examples/jax/encoder/test_single_gpu_fp8_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,99 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with FP8 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from cuda import cudart
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te
from transformer_engine.jax.fp8 import FP8Helper
from transformer_engine.common.recipe import Format as FP8Format
from transformer_engine.common.recipe import DelayedScaling

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def gpu_has_fp8():
"""GPU arch has to support FP8"""
cudaSuccess = cudart.cudaError_t.cudaSuccess
ret, gpu_id = cudart.cudaGetDevice()
assert ret == cudaSuccess
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMajor
_, major = cudart.cudaDeviceGetAttribute(flag, gpu_id)
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMinor
_, minor = cudart.cudaDeviceGetAttribute(flag, gpu_id)
sm_arch = major * 10 + minor
return sm_arch >= 89


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
others = FP8Helper.update_fp8_metas(grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
if gpu_has_fp8() is False:
print("GPU doesn't support FP8")
return

rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)
optimizer = optax.sgd(0.001, 0.9)

with te.fp8_autocast(enabled=True, fp8_recipe=DelayedScaling(fp8_format=FP8Format.HYBRID)):
encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)
Comment thread
timmoon10 marked this conversation as resolved.
assert "fp8" in str(jax.make_jaxpr(jitted_train_step)(inputs, state, variables))

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
1 change: 1 addition & 0 deletions qa/L0_jax_unittest/test.sh
Original file line numberDiff line numberDiff line change
Expand Up@@ -6,3 +6,4 @@ set -xe

: ${TE_PATH:=/opt/transformerengine}
pytest -Wignore -v $TE_PATH/tests/jax
pytest -Wignore -v $TE_PATH/examples/jax
38 changes: 17 additions & 21 deletions tests/jax/test_custom_call_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -60,15 +60,15 @@ def test_compile_bf16(self):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
# x = input, matrix 2d
# y = input, matrix 2d (weight)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -84,13 +84,13 @@ def test_compile_fp8(self, compute_type):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -104,13 +104,13 @@ def test_forward_bf16(self, m, n, k):
b = jax.random.normal(subkeys[1], (k, n), jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None))
primitive_out = fp8_dot(fp8_gemm_pkg, *_format2dtypes(None))
ref_out = jnp.dot(a, b)

assert_allclose(primitive_out, ref_out)
Expand All@@ -128,20 +128,20 @@ def test_forward_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

# calculate amax
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
# calculate scale by amax
fp8_meta = FP8Helper._update_fp8_metas_impl(fp8_meta)

fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
ref_out = jnp.dot(a, b)

ref_out = ref_out.astype(jnp.float32)
Expand All@@ -158,13 +158,13 @@ def test_grad_bf16(self, m, n, k):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.mean(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.mean(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

def ref_func(x, y):
return jnp.mean(jnp.dot(x, y))
Expand DownExpand Up@@ -193,15 +193,15 @@ def test_grad_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

def primitive_func(x, y, metas):
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], *metas)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

def ref_func(x, y):
return jnp.sum(jnp.dot(x, y))
Expand DownExpand Up@@ -232,13 +232,13 @@ def test_contracting_dims_bf16(self):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None), ((2, 3), (0, 1))))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None), ((2, 3), (0, 1))))

def ref_func(x, y):
return jnp.sum(lax.dot_general(x, y, dimension_numbers=(((2, 3), (0, 1)), ((), ()))))
Expand DownExpand Up@@ -266,7 +266,7 @@ def test_grad_fp8_mlp_randint(self, m, n, k):
s = jax.random.uniform(subkeys[3], (k,), jnp.bfloat16, 5, 8)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM * 2)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
Expand All@@ -283,7 +283,6 @@ def primitive_func(x, ln_s, y, z, metas):
ln_s,
None,
"rmsnorm",
0,
*compute_type,
activations=activations))

Expand All@@ -305,7 +304,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
amax: jnp.ndarray,
scale: jnp.ndarray,
scale_inv: jnp.ndarray,
amax_history_idx: int,
fwd_dtype,
bwd_dtype,
epsilon=1e-6,
Expand All@@ -323,7 +321,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[:FP8Helper.NUM_META_PER_GEMM],
scale_inv[:FP8Helper.NUM_META_PER_GEMM])
linear_1_out = fp8_dot(fp8_gemm_1_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -341,7 +338,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[FP8Helper.NUM_META_PER_GEMM:],
scale_inv[FP8Helper.NUM_META_PER_GEMM:])
output = fp8_dot(fp8_gemm_2_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -350,7 +346,7 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,

def ref_func(x, ln_s, y, z, metas):
return jnp.mean(
fp8_ln_mlp_py(x, ln_s, y, z, *metas, 0, *compute_type, activations=activations))
fp8_ln_mlp_py(x, ln_s, y, z, *metas, *compute_type, activations=activations))

value_n_grad_primitive_func = jit(value_and_grad(primitive_func, (0, 1, 2, 3)))
value_n_grad_ref_func = jit(value_and_grad(ref_func, (0, 1, 2, 3)))
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 75 additions & 0 deletions examples/jax/encoder/test_single_gpu_bf16_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,75 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with BF16 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)

encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
optimizer = optax.sgd(0.001, 0.9)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
99 changes: 99 additions & 0 deletions examples/jax/encoder/test_single_gpu_fp8_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,99 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with FP8 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from cuda import cudart
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te
from transformer_engine.jax.fp8 import FP8Helper
from transformer_engine.common.recipe import Format as FP8Format
from transformer_engine.common.recipe import DelayedScaling

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def gpu_has_fp8():
"""GPU arch has to support FP8"""
cudaSuccess = cudart.cudaError_t.cudaSuccess
ret, gpu_id = cudart.cudaGetDevice()
assert ret == cudaSuccess
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMajor
_, major = cudart.cudaDeviceGetAttribute(flag, gpu_id)
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMinor
_, minor = cudart.cudaDeviceGetAttribute(flag, gpu_id)
sm_arch = major * 10 + minor
return sm_arch >= 89


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
others = FP8Helper.update_fp8_metas(grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
if gpu_has_fp8() is False:
print("GPU doesn't support FP8")
return

rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)
optimizer = optax.sgd(0.001, 0.9)

with te.fp8_autocast(enabled=True, fp8_recipe=DelayedScaling(fp8_format=FP8Format.HYBRID)):
encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)
Comment thread
timmoon10 marked this conversation as resolved.
assert "fp8" in str(jax.make_jaxpr(jitted_train_step)(inputs, state, variables))

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
1 change: 1 addition & 0 deletions qa/L0_jax_unittest/test.sh
Original file line numberDiff line numberDiff line change
Expand Up@@ -6,3 +6,4 @@ set -xe

: ${TE_PATH:=/opt/transformerengine}
pytest -Wignore -v $TE_PATH/tests/jax
pytest -Wignore -v $TE_PATH/examples/jax
38 changes: 17 additions & 21 deletions tests/jax/test_custom_call_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -60,15 +60,15 @@ def test_compile_bf16(self):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
# x = input, matrix 2d
# y = input, matrix 2d (weight)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -84,13 +84,13 @@ def test_compile_fp8(self, compute_type):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -104,13 +104,13 @@ def test_forward_bf16(self, m, n, k):
b = jax.random.normal(subkeys[1], (k, n), jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None))
primitive_out = fp8_dot(fp8_gemm_pkg, *_format2dtypes(None))
ref_out = jnp.dot(a, b)

assert_allclose(primitive_out, ref_out)
Expand All@@ -128,20 +128,20 @@ def test_forward_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

# calculate amax
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
# calculate scale by amax
fp8_meta = FP8Helper._update_fp8_metas_impl(fp8_meta)

fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
ref_out = jnp.dot(a, b)

ref_out = ref_out.astype(jnp.float32)
Expand All@@ -158,13 +158,13 @@ def test_grad_bf16(self, m, n, k):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.mean(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.mean(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

def ref_func(x, y):
return jnp.mean(jnp.dot(x, y))
Expand DownExpand Up@@ -193,15 +193,15 @@ def test_grad_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

def primitive_func(x, y, metas):
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], *metas)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

def ref_func(x, y):
return jnp.sum(jnp.dot(x, y))
Expand DownExpand Up@@ -232,13 +232,13 @@ def test_contracting_dims_bf16(self):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None), ((2, 3), (0, 1))))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None), ((2, 3), (0, 1))))

def ref_func(x, y):
return jnp.sum(lax.dot_general(x, y, dimension_numbers=(((2, 3), (0, 1)), ((), ()))))
Expand DownExpand Up@@ -266,7 +266,7 @@ def test_grad_fp8_mlp_randint(self, m, n, k):
s = jax.random.uniform(subkeys[3], (k,), jnp.bfloat16, 5, 8)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM * 2)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
Expand All@@ -283,7 +283,6 @@ def primitive_func(x, ln_s, y, z, metas):
ln_s,
None,
"rmsnorm",
0,
*compute_type,
activations=activations))

Expand All@@ -305,7 +304,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
amax: jnp.ndarray,
scale: jnp.ndarray,
scale_inv: jnp.ndarray,
amax_history_idx: int,
fwd_dtype,
bwd_dtype,
epsilon=1e-6,
Expand All@@ -323,7 +321,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[:FP8Helper.NUM_META_PER_GEMM],
scale_inv[:FP8Helper.NUM_META_PER_GEMM])
linear_1_out = fp8_dot(fp8_gemm_1_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -341,7 +338,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[FP8Helper.NUM_META_PER_GEMM:],
scale_inv[FP8Helper.NUM_META_PER_GEMM:])
output = fp8_dot(fp8_gemm_2_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -350,7 +346,7 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,

def ref_func(x, ln_s, y, z, metas):
return jnp.mean(
fp8_ln_mlp_py(x, ln_s, y, z, *metas, 0, *compute_type, activations=activations))
fp8_ln_mlp_py(x, ln_s, y, z, *metas, *compute_type, activations=activations))

value_n_grad_primitive_func = jit(value_and_grad(primitive_func, (0, 1, 2, 3)))
value_n_grad_ref_func = jit(value_and_grad(ref_func, (0, 1, 2, 3)))
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length > 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 75 additions & 0 deletions examples/jax/encoder/test_single_gpu_bf16_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,75 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with BF16 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)

encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
optimizer = optax.sgd(0.001, 0.9)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
99 changes: 99 additions & 0 deletions examples/jax/encoder/test_single_gpu_fp8_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,99 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with FP8 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from cuda import cudart
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te
from transformer_engine.jax.fp8 import FP8Helper
from transformer_engine.common.recipe import Format as FP8Format
from transformer_engine.common.recipe import DelayedScaling

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def gpu_has_fp8():
"""GPU arch has to support FP8"""
cudaSuccess = cudart.cudaError_t.cudaSuccess
ret, gpu_id = cudart.cudaGetDevice()
assert ret == cudaSuccess
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMajor
_, major = cudart.cudaDeviceGetAttribute(flag, gpu_id)
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMinor
_, minor = cudart.cudaDeviceGetAttribute(flag, gpu_id)
sm_arch = major * 10 + minor
return sm_arch >= 89


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
others = FP8Helper.update_fp8_metas(grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
if gpu_has_fp8() is False:
print("GPU doesn't support FP8")
return

rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)
optimizer = optax.sgd(0.001, 0.9)

with te.fp8_autocast(enabled=True, fp8_recipe=DelayedScaling(fp8_format=FP8Format.HYBRID)):
encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)
Comment thread
timmoon10 marked this conversation as resolved.
assert "fp8" in str(jax.make_jaxpr(jitted_train_step)(inputs, state, variables))

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
1 change: 1 addition & 0 deletions qa/L0_jax_unittest/test.sh
Original file line numberDiff line numberDiff line change
Expand Up@@ -6,3 +6,4 @@ set -xe

: ${TE_PATH:=/opt/transformerengine}
pytest -Wignore -v $TE_PATH/tests/jax
pytest -Wignore -v $TE_PATH/examples/jax
38 changes: 17 additions & 21 deletions tests/jax/test_custom_call_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -60,15 +60,15 @@ def test_compile_bf16(self):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
# x = input, matrix 2d
# y = input, matrix 2d (weight)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -84,13 +84,13 @@ def test_compile_fp8(self, compute_type):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -104,13 +104,13 @@ def test_forward_bf16(self, m, n, k):
b = jax.random.normal(subkeys[1], (k, n), jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None))
primitive_out = fp8_dot(fp8_gemm_pkg, *_format2dtypes(None))
ref_out = jnp.dot(a, b)

assert_allclose(primitive_out, ref_out)
Expand All@@ -128,20 +128,20 @@ def test_forward_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

# calculate amax
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
# calculate scale by amax
fp8_meta = FP8Helper._update_fp8_metas_impl(fp8_meta)

fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
ref_out = jnp.dot(a, b)

ref_out = ref_out.astype(jnp.float32)
Expand All@@ -158,13 +158,13 @@ def test_grad_bf16(self, m, n, k):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.mean(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.mean(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

def ref_func(x, y):
return jnp.mean(jnp.dot(x, y))
Expand DownExpand Up@@ -193,15 +193,15 @@ def test_grad_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

def primitive_func(x, y, metas):
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], *metas)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

def ref_func(x, y):
return jnp.sum(jnp.dot(x, y))
Expand DownExpand Up@@ -232,13 +232,13 @@ def test_contracting_dims_bf16(self):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None), ((2, 3), (0, 1))))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None), ((2, 3), (0, 1))))

def ref_func(x, y):
return jnp.sum(lax.dot_general(x, y, dimension_numbers=(((2, 3), (0, 1)), ((), ()))))
Expand DownExpand Up@@ -266,7 +266,7 @@ def test_grad_fp8_mlp_randint(self, m, n, k):
s = jax.random.uniform(subkeys[3], (k,), jnp.bfloat16, 5, 8)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM * 2)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
Expand All@@ -283,7 +283,6 @@ def primitive_func(x, ln_s, y, z, metas):
ln_s,
None,
"rmsnorm",
0,
*compute_type,
activations=activations))

Expand All@@ -305,7 +304,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
amax: jnp.ndarray,
scale: jnp.ndarray,
scale_inv: jnp.ndarray,
amax_history_idx: int,
fwd_dtype,
bwd_dtype,
epsilon=1e-6,
Expand All@@ -323,7 +321,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[:FP8Helper.NUM_META_PER_GEMM],
scale_inv[:FP8Helper.NUM_META_PER_GEMM])
linear_1_out = fp8_dot(fp8_gemm_1_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -341,7 +338,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[FP8Helper.NUM_META_PER_GEMM:],
scale_inv[FP8Helper.NUM_META_PER_GEMM:])
output = fp8_dot(fp8_gemm_2_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -350,7 +346,7 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,

def ref_func(x, ln_s, y, z, metas):
return jnp.mean(
fp8_ln_mlp_py(x, ln_s, y, z, *metas, 0, *compute_type, activations=activations))
fp8_ln_mlp_py(x, ln_s, y, z, *metas, *compute_type, activations=activations))

value_n_grad_primitive_func = jit(value_and_grad(primitive_func, (0, 1, 2, 3)))
value_n_grad_ref_func = jit(value_and_grad(ref_func, (0, 1, 2, 3)))
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 75 additions & 0 deletions examples/jax/encoder/test_single_gpu_bf16_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,75 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with BF16 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)

encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
optimizer = optax.sgd(0.001, 0.9)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
99 changes: 99 additions & 0 deletions examples/jax/encoder/test_single_gpu_fp8_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,99 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with FP8 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from cuda import cudart
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te
from transformer_engine.jax.fp8 import FP8Helper
from transformer_engine.common.recipe import Format as FP8Format
from transformer_engine.common.recipe import DelayedScaling

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def gpu_has_fp8():
"""GPU arch has to support FP8"""
cudaSuccess = cudart.cudaError_t.cudaSuccess
ret, gpu_id = cudart.cudaGetDevice()
assert ret == cudaSuccess
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMajor
_, major = cudart.cudaDeviceGetAttribute(flag, gpu_id)
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMinor
_, minor = cudart.cudaDeviceGetAttribute(flag, gpu_id)
sm_arch = major * 10 + minor
return sm_arch >= 89


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
others = FP8Helper.update_fp8_metas(grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
if gpu_has_fp8() is False:
print("GPU doesn't support FP8")
return

rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)
optimizer = optax.sgd(0.001, 0.9)

with te.fp8_autocast(enabled=True, fp8_recipe=DelayedScaling(fp8_format=FP8Format.HYBRID)):
encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)
Comment thread
timmoon10 marked this conversation as resolved.
assert "fp8" in str(jax.make_jaxpr(jitted_train_step)(inputs, state, variables))

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
1 change: 1 addition & 0 deletions qa/L0_jax_unittest/test.sh
Original file line numberDiff line numberDiff line change
Expand Up@@ -6,3 +6,4 @@ set -xe

: ${TE_PATH:=/opt/transformerengine}
pytest -Wignore -v $TE_PATH/tests/jax
pytest -Wignore -v $TE_PATH/examples/jax
38 changes: 17 additions & 21 deletions tests/jax/test_custom_call_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -60,15 +60,15 @@ def test_compile_bf16(self):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
# x = input, matrix 2d
# y = input, matrix 2d (weight)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -84,13 +84,13 @@ def test_compile_fp8(self, compute_type):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -104,13 +104,13 @@ def test_forward_bf16(self, m, n, k):
b = jax.random.normal(subkeys[1], (k, n), jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None))
primitive_out = fp8_dot(fp8_gemm_pkg, *_format2dtypes(None))
ref_out = jnp.dot(a, b)

assert_allclose(primitive_out, ref_out)
Expand All@@ -128,20 +128,20 @@ def test_forward_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

# calculate amax
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
# calculate scale by amax
fp8_meta = FP8Helper._update_fp8_metas_impl(fp8_meta)

fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
ref_out = jnp.dot(a, b)

ref_out = ref_out.astype(jnp.float32)
Expand All@@ -158,13 +158,13 @@ def test_grad_bf16(self, m, n, k):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.mean(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.mean(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

def ref_func(x, y):
return jnp.mean(jnp.dot(x, y))
Expand DownExpand Up@@ -193,15 +193,15 @@ def test_grad_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

def primitive_func(x, y, metas):
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], *metas)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

def ref_func(x, y):
return jnp.sum(jnp.dot(x, y))
Expand DownExpand Up@@ -232,13 +232,13 @@ def test_contracting_dims_bf16(self):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None), ((2, 3), (0, 1))))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None), ((2, 3), (0, 1))))

def ref_func(x, y):
return jnp.sum(lax.dot_general(x, y, dimension_numbers=(((2, 3), (0, 1)), ((), ()))))
Expand DownExpand Up@@ -266,7 +266,7 @@ def test_grad_fp8_mlp_randint(self, m, n, k):
s = jax.random.uniform(subkeys[3], (k,), jnp.bfloat16, 5, 8)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM * 2)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
Expand All@@ -283,7 +283,6 @@ def primitive_func(x, ln_s, y, z, metas):
ln_s,
None,
"rmsnorm",
0,
*compute_type,
activations=activations))

Expand All@@ -305,7 +304,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
amax: jnp.ndarray,
scale: jnp.ndarray,
scale_inv: jnp.ndarray,
amax_history_idx: int,
fwd_dtype,
bwd_dtype,
epsilon=1e-6,
Expand All@@ -323,7 +321,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[:FP8Helper.NUM_META_PER_GEMM],
scale_inv[:FP8Helper.NUM_META_PER_GEMM])
linear_1_out = fp8_dot(fp8_gemm_1_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -341,7 +338,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[FP8Helper.NUM_META_PER_GEMM:],
scale_inv[FP8Helper.NUM_META_PER_GEMM:])
output = fp8_dot(fp8_gemm_2_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -350,7 +346,7 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,

def ref_func(x, ln_s, y, z, metas):
return jnp.mean(
fp8_ln_mlp_py(x, ln_s, y, z, *metas, 0, *compute_type, activations=activations))
fp8_ln_mlp_py(x, ln_s, y, z, *metas, *compute_type, activations=activations))

value_n_grad_primitive_func = jit(value_and_grad(primitive_func, (0, 1, 2, 3)))
value_n_grad_ref_func = jit(value_and_grad(ref_func, (0, 1, 2, 3)))
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 75 additions & 0 deletions examples/jax/encoder/test_single_gpu_bf16_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,75 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with BF16 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)

encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
optimizer = optax.sgd(0.001, 0.9)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
99 changes: 99 additions & 0 deletions examples/jax/encoder/test_single_gpu_fp8_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,99 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with FP8 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from cuda import cudart
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te
from transformer_engine.jax.fp8 import FP8Helper
from transformer_engine.common.recipe import Format as FP8Format
from transformer_engine.common.recipe import DelayedScaling

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def gpu_has_fp8():
"""GPU arch has to support FP8"""
cudaSuccess = cudart.cudaError_t.cudaSuccess
ret, gpu_id = cudart.cudaGetDevice()
assert ret == cudaSuccess
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMajor
_, major = cudart.cudaDeviceGetAttribute(flag, gpu_id)
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMinor
_, minor = cudart.cudaDeviceGetAttribute(flag, gpu_id)
sm_arch = major * 10 + minor
return sm_arch >= 89


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
others = FP8Helper.update_fp8_metas(grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
if gpu_has_fp8() is False:
print("GPU doesn't support FP8")
return

rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)
optimizer = optax.sgd(0.001, 0.9)

with te.fp8_autocast(enabled=True, fp8_recipe=DelayedScaling(fp8_format=FP8Format.HYBRID)):
encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)
Comment thread
timmoon10 marked this conversation as resolved.
assert "fp8" in str(jax.make_jaxpr(jitted_train_step)(inputs, state, variables))

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
1 change: 1 addition & 0 deletions qa/L0_jax_unittest/test.sh
Original file line numberDiff line numberDiff line change
Expand Up@@ -6,3 +6,4 @@ set -xe

: ${TE_PATH:=/opt/transformerengine}
pytest -Wignore -v $TE_PATH/tests/jax
pytest -Wignore -v $TE_PATH/examples/jax
38 changes: 17 additions & 21 deletions tests/jax/test_custom_call_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -60,15 +60,15 @@ def test_compile_bf16(self):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
# x = input, matrix 2d
# y = input, matrix 2d (weight)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -84,13 +84,13 @@ def test_compile_fp8(self, compute_type):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -104,13 +104,13 @@ def test_forward_bf16(self, m, n, k):
b = jax.random.normal(subkeys[1], (k, n), jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None))
primitive_out = fp8_dot(fp8_gemm_pkg, *_format2dtypes(None))
ref_out = jnp.dot(a, b)

assert_allclose(primitive_out, ref_out)
Expand All@@ -128,20 +128,20 @@ def test_forward_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

# calculate amax
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
# calculate scale by amax
fp8_meta = FP8Helper._update_fp8_metas_impl(fp8_meta)

fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
ref_out = jnp.dot(a, b)

ref_out = ref_out.astype(jnp.float32)
Expand All@@ -158,13 +158,13 @@ def test_grad_bf16(self, m, n, k):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.mean(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.mean(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

def ref_func(x, y):
return jnp.mean(jnp.dot(x, y))
Expand DownExpand Up@@ -193,15 +193,15 @@ def test_grad_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

def primitive_func(x, y, metas):
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], *metas)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

def ref_func(x, y):
return jnp.sum(jnp.dot(x, y))
Expand DownExpand Up@@ -232,13 +232,13 @@ def test_contracting_dims_bf16(self):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None), ((2, 3), (0, 1))))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None), ((2, 3), (0, 1))))

def ref_func(x, y):
return jnp.sum(lax.dot_general(x, y, dimension_numbers=(((2, 3), (0, 1)), ((), ()))))
Expand DownExpand Up@@ -266,7 +266,7 @@ def test_grad_fp8_mlp_randint(self, m, n, k):
s = jax.random.uniform(subkeys[3], (k,), jnp.bfloat16, 5, 8)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM * 2)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
Expand All@@ -283,7 +283,6 @@ def primitive_func(x, ln_s, y, z, metas):
ln_s,
None,
"rmsnorm",
0,
*compute_type,
activations=activations))

Expand All@@ -305,7 +304,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
amax: jnp.ndarray,
scale: jnp.ndarray,
scale_inv: jnp.ndarray,
amax_history_idx: int,
fwd_dtype,
bwd_dtype,
epsilon=1e-6,
Expand All@@ -323,7 +321,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[:FP8Helper.NUM_META_PER_GEMM],
scale_inv[:FP8Helper.NUM_META_PER_GEMM])
linear_1_out = fp8_dot(fp8_gemm_1_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -341,7 +338,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[FP8Helper.NUM_META_PER_GEMM:],
scale_inv[FP8Helper.NUM_META_PER_GEMM:])
output = fp8_dot(fp8_gemm_2_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -350,7 +346,7 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,

def ref_func(x, ln_s, y, z, metas):
return jnp.mean(
fp8_ln_mlp_py(x, ln_s, y, z, *metas, 0, *compute_type, activations=activations))
fp8_ln_mlp_py(x, ln_s, y, z, *metas, *compute_type, activations=activations))

value_n_grad_primitive_func = jit(value_and_grad(primitive_func, (0, 1, 2, 3)))
value_n_grad_ref_func = jit(value_and_grad(ref_func, (0, 1, 2, 3)))
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 75 additions & 0 deletions examples/jax/encoder/test_single_gpu_bf16_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,75 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with BF16 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)

encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
optimizer = optax.sgd(0.001, 0.9)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
99 changes: 99 additions & 0 deletions examples/jax/encoder/test_single_gpu_fp8_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,99 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with FP8 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from cuda import cudart
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te
from transformer_engine.jax.fp8 import FP8Helper
from transformer_engine.common.recipe import Format as FP8Format
from transformer_engine.common.recipe import DelayedScaling

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def gpu_has_fp8():
"""GPU arch has to support FP8"""
cudaSuccess = cudart.cudaError_t.cudaSuccess
ret, gpu_id = cudart.cudaGetDevice()
assert ret == cudaSuccess
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMajor
_, major = cudart.cudaDeviceGetAttribute(flag, gpu_id)
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMinor
_, minor = cudart.cudaDeviceGetAttribute(flag, gpu_id)
sm_arch = major * 10 + minor
return sm_arch >= 89


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
others = FP8Helper.update_fp8_metas(grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
if gpu_has_fp8() is False:
print("GPU doesn't support FP8")
return

rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)
optimizer = optax.sgd(0.001, 0.9)

with te.fp8_autocast(enabled=True, fp8_recipe=DelayedScaling(fp8_format=FP8Format.HYBRID)):
encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)
Comment thread
timmoon10 marked this conversation as resolved.
assert "fp8" in str(jax.make_jaxpr(jitted_train_step)(inputs, state, variables))

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
1 change: 1 addition & 0 deletions qa/L0_jax_unittest/test.sh
Original file line numberDiff line numberDiff line change
Expand Up@@ -6,3 +6,4 @@ set -xe

: ${TE_PATH:=/opt/transformerengine}
pytest -Wignore -v $TE_PATH/tests/jax
pytest -Wignore -v $TE_PATH/examples/jax
38 changes: 17 additions & 21 deletions tests/jax/test_custom_call_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -60,15 +60,15 @@ def test_compile_bf16(self):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
# x = input, matrix 2d
# y = input, matrix 2d (weight)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -84,13 +84,13 @@ def test_compile_fp8(self, compute_type):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -104,13 +104,13 @@ def test_forward_bf16(self, m, n, k):
b = jax.random.normal(subkeys[1], (k, n), jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None))
primitive_out = fp8_dot(fp8_gemm_pkg, *_format2dtypes(None))
ref_out = jnp.dot(a, b)

assert_allclose(primitive_out, ref_out)
Expand All@@ -128,20 +128,20 @@ def test_forward_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

# calculate amax
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
# calculate scale by amax
fp8_meta = FP8Helper._update_fp8_metas_impl(fp8_meta)

fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
ref_out = jnp.dot(a, b)

ref_out = ref_out.astype(jnp.float32)
Expand All@@ -158,13 +158,13 @@ def test_grad_bf16(self, m, n, k):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.mean(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.mean(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

def ref_func(x, y):
return jnp.mean(jnp.dot(x, y))
Expand DownExpand Up@@ -193,15 +193,15 @@ def test_grad_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

def primitive_func(x, y, metas):
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], *metas)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

def ref_func(x, y):
return jnp.sum(jnp.dot(x, y))
Expand DownExpand Up@@ -232,13 +232,13 @@ def test_contracting_dims_bf16(self):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None), ((2, 3), (0, 1))))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None), ((2, 3), (0, 1))))

def ref_func(x, y):
return jnp.sum(lax.dot_general(x, y, dimension_numbers=(((2, 3), (0, 1)), ((), ()))))
Expand DownExpand Up@@ -266,7 +266,7 @@ def test_grad_fp8_mlp_randint(self, m, n, k):
s = jax.random.uniform(subkeys[3], (k,), jnp.bfloat16, 5, 8)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM * 2)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
Expand All@@ -283,7 +283,6 @@ def primitive_func(x, ln_s, y, z, metas):
ln_s,
None,
"rmsnorm",
0,
*compute_type,
activations=activations))

Expand All@@ -305,7 +304,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
amax: jnp.ndarray,
scale: jnp.ndarray,
scale_inv: jnp.ndarray,
amax_history_idx: int,
fwd_dtype,
bwd_dtype,
epsilon=1e-6,
Expand All@@ -323,7 +321,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[:FP8Helper.NUM_META_PER_GEMM],
scale_inv[:FP8Helper.NUM_META_PER_GEMM])
linear_1_out = fp8_dot(fp8_gemm_1_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -341,7 +338,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[FP8Helper.NUM_META_PER_GEMM:],
scale_inv[FP8Helper.NUM_META_PER_GEMM:])
output = fp8_dot(fp8_gemm_2_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -350,7 +346,7 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,

def ref_func(x, ln_s, y, z, metas):
return jnp.mean(
fp8_ln_mlp_py(x, ln_s, y, z, *metas, 0, *compute_type, activations=activations))
fp8_ln_mlp_py(x, ln_s, y, z, *metas, *compute_type, activations=activations))

value_n_grad_primitive_func = jit(value_and_grad(primitive_func, (0, 1, 2, 3)))
value_n_grad_ref_func = jit(value_and_grad(ref_func, (0, 1, 2, 3)))
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 75 additions & 0 deletions examples/jax/encoder/test_single_gpu_bf16_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,75 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with BF16 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)

encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
optimizer = optax.sgd(0.001, 0.9)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
99 changes: 99 additions & 0 deletions examples/jax/encoder/test_single_gpu_fp8_training.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,99 @@
# Copyright (c) 2022-2023, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
""" Encoder with FP8 Training on single GPU"""
import jax
import jax.numpy as jnp
import optax
from cuda import cudart
from flax.core.frozen_dict import FrozenDict
from flax.training import train_state

import transformer_engine.jax as te
from transformer_engine.jax.fp8 import FP8Helper
from transformer_engine.common.recipe import Format as FP8Format
from transformer_engine.common.recipe import DelayedScaling

PARAMS_KEY = 'params'

BATCH = 32
SEQLEN = 512
HIDDEN = 1024


def gpu_has_fp8():
"""GPU arch has to support FP8"""
cudaSuccess = cudart.cudaError_t.cudaSuccess
ret, gpu_id = cudart.cudaGetDevice()
assert ret == cudaSuccess
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMajor
_, major = cudart.cudaDeviceGetAttribute(flag, gpu_id)
flag = cudart.cudaDeviceAttr.cudaDevAttrComputeCapabilityMinor
_, minor = cudart.cudaDeviceGetAttribute(flag, gpu_id)
sm_arch = major * 10 + minor
return sm_arch >= 89


def network():
"""NLP Encoder"""
encoder = te.TransformerLayer(hidden_size=HIDDEN,
mlp_hidden_size=4 * HIDDEN,
hidden_dropout=0.0,
attention_dropout=0.0,
layernorm_type='rmsnorm',
mlp_activations=('gelu', 'linear'),
layer_type=te.TransformerLayerType.ENCODER,
transpose_batch_sequence=True,
dtype=jnp.bfloat16)
return encoder


def synthesis_data(data_rng):
"""Dataset generator"""
return jax.random.normal(data_rng, [SEQLEN, BATCH, HIDDEN], jnp.bfloat16)


def train_step(batch, state, others):
"""Training function."""

def loss_fn(collections):
logits = state.apply_fn(collections, batch)
loss = jnp.mean(logits)
return loss

grad_fn = jax.value_and_grad(loss_fn)
loss, grads = grad_fn(FrozenDict({PARAMS_KEY: state.params, **others}))
grads, params_grads = grads.pop(PARAMS_KEY)
state = state.apply_gradients(grads=params_grads)
others = FP8Helper.update_fp8_metas(grads)
return loss, state, others


def test_encoder():
"""Encoder example"""
if gpu_has_fp8() is False:
print("GPU doesn't support FP8")
return

rng = jax.random.PRNGKey(0)
rng, init_rng, data_rng = jax.random.split(rng, 3)
inputs = synthesis_data(data_rng)
optimizer = optax.sgd(0.001, 0.9)

with te.fp8_autocast(enabled=True, fp8_recipe=DelayedScaling(fp8_format=FP8Format.HYBRID)):
encoder = network()
variables = jax.jit(encoder.init)(init_rng, inputs)
variables, params = variables.pop(PARAMS_KEY)
state = train_state.TrainState.create(apply_fn=encoder.apply, params=params, tx=optimizer)
jitted_train_step = jax.jit(train_step)
Comment thread
timmoon10 marked this conversation as resolved.
assert "fp8" in str(jax.make_jaxpr(jitted_train_step)(inputs, state, variables))

for i in range(5):
rng, data_rng = jax.random.split(rng)
inputs = synthesis_data(data_rng)
loss, state, variables = jitted_train_step(inputs, state, variables)
print(f"Step {i} - Loss: {loss}")


if __name__ == "__main__":
test_encoder()
1 change: 1 addition & 0 deletions qa/L0_jax_unittest/test.sh
Original file line numberDiff line numberDiff line change
Expand Up@@ -6,3 +6,4 @@ set -xe

: ${TE_PATH:=/opt/transformerengine}
pytest -Wignore -v $TE_PATH/tests/jax
pytest -Wignore -v $TE_PATH/examples/jax
38 changes: 17 additions & 21 deletions tests/jax/test_custom_call_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -60,15 +60,15 @@ def test_compile_bf16(self):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
# x = input, matrix 2d
# y = input, matrix 2d (weight)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -84,13 +84,13 @@ def test_compile_fp8(self, compute_type):

def func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

value_n_grad_func = value_and_grad(func, (0, 1))
value_n_grad_func_compiled = jit(value_n_grad_func).lower(a, b).compile()
Expand All@@ -104,13 +104,13 @@ def test_forward_bf16(self, m, n, k):
b = jax.random.normal(subkeys[1], (k, n), jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None))
primitive_out = fp8_dot(fp8_gemm_pkg, *_format2dtypes(None))
ref_out = jnp.dot(a, b)

assert_allclose(primitive_out, ref_out)
Expand All@@ -128,20 +128,20 @@ def test_forward_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

# calculate amax
fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
# calculate scale by amax
fp8_meta = FP8Helper._update_fp8_metas_impl(fp8_meta)

fp8_gemm_pkg = FP8GemmPackage(1, a, [b], *fp8_meta)
primitive_out = fp8_dot(fp8_gemm_pkg, 0, *compute_type)
primitive_out = fp8_dot(fp8_gemm_pkg, *compute_type)
ref_out = jnp.dot(a, b)

ref_out = ref_out.astype(jnp.float32)
Expand All@@ -158,13 +158,13 @@ def test_grad_bf16(self, m, n, k):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.mean(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None)))
return jnp.mean(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None)))

def ref_func(x, y):
return jnp.mean(jnp.dot(x, y))
Expand DownExpand Up@@ -193,15 +193,15 @@ def test_grad_fp8_randint(self, m, n, k, compute_type):
b = jax.random.randint(subkeys[1], (k, n), min_val, max_val).astype(jnp.bfloat16)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_meta = [fp8_max, fp8_metas_amax, fp8_metas_scale, fp8_metas_scale_inv]

def primitive_func(x, y, metas):
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], *metas)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *compute_type))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *compute_type))

def ref_func(x, y):
return jnp.sum(jnp.dot(x, y))
Expand DownExpand Up@@ -232,13 +232,13 @@ def test_contracting_dims_bf16(self):

def primitive_func(x, y):
fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM, 1), jnp.float32)
fp8_gemm_pkg = FP8GemmPackage(1, x, [y], fp8_max, fp8_metas_amax, fp8_metas_scale,
fp8_metas_scale_inv)
return jnp.sum(fp8_dot(fp8_gemm_pkg, 0, *_format2dtypes(None), ((2, 3), (0, 1))))
return jnp.sum(fp8_dot(fp8_gemm_pkg, *_format2dtypes(None), ((2, 3), (0, 1))))

def ref_func(x, y):
return jnp.sum(lax.dot_general(x, y, dimension_numbers=(((2, 3), (0, 1)), ((), ()))))
Expand DownExpand Up@@ -266,7 +266,7 @@ def test_grad_fp8_mlp_randint(self, m, n, k):
s = jax.random.uniform(subkeys[3], (k,), jnp.bfloat16, 5, 8)

fp8_max = FP8Helper.generate_fp8_max_array(FP8Helper.NUM_META_PER_GEMM * 2)
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_SIZE),
fp8_metas_amax = jnp.zeros((FP8Helper.NUM_META_PER_GEMM * 2, FP8Helper.AMAX_HISTORY_LEN),
jnp.float32)
fp8_metas_scale = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
fp8_metas_scale_inv = jnp.ones((FP8Helper.NUM_META_PER_GEMM * 2, 1), jnp.float32)
Expand All@@ -283,7 +283,6 @@ def primitive_func(x, ln_s, y, z, metas):
ln_s,
None,
"rmsnorm",
0,
*compute_type,
activations=activations))

Expand All@@ -305,7 +304,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
amax: jnp.ndarray,
scale: jnp.ndarray,
scale_inv: jnp.ndarray,
amax_history_idx: int,
fwd_dtype,
bwd_dtype,
epsilon=1e-6,
Expand All@@ -323,7 +321,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[:FP8Helper.NUM_META_PER_GEMM],
scale_inv[:FP8Helper.NUM_META_PER_GEMM])
linear_1_out = fp8_dot(fp8_gemm_1_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -341,7 +338,6 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,
scale[FP8Helper.NUM_META_PER_GEMM:],
scale_inv[FP8Helper.NUM_META_PER_GEMM:])
output = fp8_dot(fp8_gemm_2_pkg,
amax_history_idx,
fwd_dtype,
bwd_dtype,
contracting_dims,
Expand All@@ -350,7 +346,7 @@ def fp8_ln_mlp_py(inputs: jnp.ndarray,

def ref_func(x, ln_s, y, z, metas):
return jnp.mean(
fp8_ln_mlp_py(x, ln_s, y, z, *metas, 0, *compute_type, activations=activations))
fp8_ln_mlp_py(x, ln_s, y, z, *metas, *compute_type, activations=activations))

value_n_grad_primitive_func = jit(value_and_grad(primitive_func, (0, 1, 2, 3)))
value_n_grad_ref_func = jit(value_and_grad(ref_func, (0, 1, 2, 3)))
Expand Down
Loading