Merged
79 changes: 48 additions & 31 deletions tests/cpp/operator/test_layernorm.cu
Original file line numberDiff line numberDiff line change
Expand Up@@ -49,14 +49,17 @@ template <typename InputType, typename OutputType>
void compute_ref_output(const InputType *data, const InputType *gamma, const InputType *beta,
OutputType *output, const float *mu, const float *rsigma,
const size_t N, const size_t H,
float *amax, float scale) {
float *amax, float scale, const bool zero_centered_gamma) {
using compute_t = float;
compute_t current_max = -1e100;
for (size_t i = 0 ; i < N; ++i) {
for (size_t j = 0; j < H; ++j) {
compute_t current = static_cast<compute_t>(data[i * H + j]);
compute_t tmp = (current - mu[i]) * rsigma[i] * static_cast<compute_t>(gamma[j]) +
static_cast<compute_t>(beta[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
compute_t tmp = (current - mu[i]) * rsigma[i] * g + static_cast<compute_t>(beta[j]);
output[i * H + j] = static_cast<OutputType>(tmp * scale);
current_max = fmaxf(current_max, fabsf(tmp));
}
Expand All@@ -70,7 +73,8 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
const InputType *gamma,
InputType *data_grad,
InputType *gamma_grad, InputType *beta_grad,
const size_t N, const size_t H) {
const size_t N, const size_t H,
const bool zero_centered_gamma) {
using compute_t = float;
std::vector<compute_t> dgamma(H, 0.f);
std::vector<compute_t> dbeta(H, 0.f);
Expand All@@ -81,7 +85,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
dgamma[j] += y * dz;
Expand All@@ -96,7 +103,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
const compute_t dx = rsigma[i] * (dy - mdyy * y - mdy);
Expand All@@ -112,7 +122,7 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
}

template <typename InputType, typename OutputType>
void performTest(const size_t N, const size_t H) {
void performTest(const size_t N, const size_t H, const bool zero_centered_gamma) {
if (sizeof(InputType) < sizeof(OutputType)) {
GTEST_SKIP() << "LN kernel does not support OutputType > InputType";
return;
Expand DownExpand Up@@ -158,32 +168,34 @@ void performTest(const size_t N, const size_t H) {

// Forward kernel
float epsilon = 1e-5;
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto fwd_function = zero_centered_gamma ? nvte_layernorm1p_fwd : nvte_layernorm_fwd;
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Backward kernel
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto bwd_function = zero_centered_gamma ? nvte_layernorm1p_bwd : nvte_layernorm_bwd;
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
dgamma_part = Tensor(dgamma_part.shape(), dgamma_part.dtype());
dbeta_part = Tensor(dbeta_part.shape(), dbeta_part.dtype());
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Reference implementations
// use the GPU stats to tighten the tolerances
Expand All@@ -201,12 +213,13 @@ void performTest(const size_t N, const size_t H) {
rsigma.cpu_dptr<float>(),
N, H,
&ref_amax,
ref_scale);
ref_scale,
zero_centered_gamma);
compute_ref_backward(dz.cpu_dptr<WeightType>(), input.cpu_dptr<InputType>(),
mu.cpu_dptr<float>(), rsigma.cpu_dptr<float>(),
gamma.cpu_dptr<WeightType>(),
ref_dx.get(), ref_dgamma.get(), ref_dbeta.get(),
N, H);
N, H, zero_centered_gamma);

cudaDeviceSynchronize();
auto err = cudaGetLastError();
Expand DownExpand Up@@ -248,7 +261,8 @@ std::vector<std::pair<size_t, size_t>> test_cases = {{2048, 12288},

class LNTestSuite : public ::testing::TestWithParam<std::tuple<transformer_engine::DType,
transformer_engine::DType,
std::pair<size_t, size_t>>> {};
std::pair<size_t, size_t>,
bool>> {};

TEST_P(LNTestSuite, TestLN) {
using namespace transformer_engine;
Expand All@@ -257,10 +271,11 @@ TEST_P(LNTestSuite, TestLN) {
const DType input_type = std::get<0>(GetParam());
const DType output_type = std::get<1>(GetParam());
const auto size = std::get<2>(GetParam());
const bool zero_centered_gamma = std::get<3>(GetParam());

TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(input_type, InputType,
TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(output_type, OutputType,
performTest<InputType, OutputType>(size.first, size.second);
performTest<InputType, OutputType>(size.first, size.second, zero_centered_gamma);
);
);
}
Expand All@@ -271,11 +286,13 @@ INSTANTIATE_TEST_SUITE_P(
::testing::Combine(
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16),
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16, DType::kFloat8E4M3),
::testing::ValuesIn(test_cases)),
::testing::ValuesIn(test_cases),
::testing::Values(false, true)),
[](const testing::TestParamInfo<LNTestSuite::ParamType>& info) {
std::string name = test::typeName(std::get<0>(info.param)) + "X" +
test::typeName(std::get<1>(info.param)) + "X" +
std::to_string(std::get<2>(info.param).first) + "X" +
std::to_string(std::get<2>(info.param).second);
std::to_string(std::get<2>(info.param).second) + "X" +
std::to_string(std::get<3>(info.param));
return name;
});
29 changes: 21 additions & 8 deletions tests/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -434,10 +434,12 @@ def forward(self, inp, weight):
@pytest.mark.parametrize("use_fp8", [False, True])
@pytest.mark.parametrize("scale_factor", [448, 112])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm(
use_fp8: bool,
scale_factor: float,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -459,7 +461,8 @@ def forward(self, inp):
inp,
self.weight,
self.bias,
self.eps)
self.eps,
zero_centered_gamma)
return ret

class TestFP8_Layernorm(nn.Module):
Expand All@@ -482,7 +485,8 @@ def forward(self, inp):
self.eps,
self.meta,
self.fp8_tensor,
self.fp8_type)
self.fp8_type,
zero_centered_gamma)

ret = cast_from_fp8(
ret,
Expand All@@ -500,7 +504,7 @@ def forward(self, inp):
do_export(model, inp, fname, use_fp8=use_fp8)
if precision not in (torch.bfloat16, ):
# TODO: FP32 has a small threshold (1e-5)
validate_result(fname, inp, model, atol=1e-3, is_fp8=use_fp8)
validate_result(fname, inp, model, atol=4e-3, is_fp8=use_fp8)


@skip_FP8
Expand DownExpand Up@@ -646,13 +650,15 @@ def forward(self, inp):
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_linear(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -676,6 +682,7 @@ def test_export_layernorm_linear(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand All@@ -698,13 +705,15 @@ def test_export_layernorm_linear(
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_mlp(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -729,6 +738,7 @@ def test_export_layernorm_mlp(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand DownExpand Up@@ -902,14 +912,16 @@ def test_export_multihead_attention(
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
use_mask: bool,
attn_mask_type: str,
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand DownExpand Up@@ -947,7 +959,8 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling).to(device='cuda')
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
validate_result(fname, inp, model, atol=1e-3)
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Add copy buttons to all
 blocks
(function() {
function addCopyButtons() {
document.querySelectorAll('pre code').forEach(function(codeBlock) {
if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;
codeBlock.parentElement.setAttribute('data-copy-added', 'true');
var btn = document.createElement('button');
btn.textContent = 'Copy';
btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';
btn.onmouseover = function() { this.style.opacity = '1'; };
btn.onmouseout = function() { this.style.opacity = '0.7'; };
btn.onclick = function() {
navigator.clipboard.writeText(codeBlock.textContent).then(function() {
btn.textContent = 'Copied!';
setTimeout(function() { btn.textContent = 'Copy'; }, 1500);
});
};
codeBlock.parentElement.style.position = 'relative';
codeBlock.parentElement.appendChild(btn);
});
}
addCopyButtons();
// Re-run on dynamic content
var observer = new MutationObserver(addCopyButtons);
observer.observe(document.body, { childList: true, subtree: true });
})();
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Skip to content
Merged
79 changes: 48 additions & 31 deletions tests/cpp/operator/test_layernorm.cu
Original file line numberDiff line numberDiff line change
Expand Up@@ -49,14 +49,17 @@ template <typename InputType, typename OutputType>
void compute_ref_output(const InputType *data, const InputType *gamma, const InputType *beta,
OutputType *output, const float *mu, const float *rsigma,
const size_t N, const size_t H,
float *amax, float scale) {
float *amax, float scale, const bool zero_centered_gamma) {
using compute_t = float;
compute_t current_max = -1e100;
for (size_t i = 0 ; i < N; ++i) {
for (size_t j = 0; j < H; ++j) {
compute_t current = static_cast<compute_t>(data[i * H + j]);
compute_t tmp = (current - mu[i]) * rsigma[i] * static_cast<compute_t>(gamma[j]) +
static_cast<compute_t>(beta[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
compute_t tmp = (current - mu[i]) * rsigma[i] * g + static_cast<compute_t>(beta[j]);
output[i * H + j] = static_cast<OutputType>(tmp * scale);
current_max = fmaxf(current_max, fabsf(tmp));
}
Expand All@@ -70,7 +73,8 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
const InputType *gamma,
InputType *data_grad,
InputType *gamma_grad, InputType *beta_grad,
const size_t N, const size_t H) {
const size_t N, const size_t H,
const bool zero_centered_gamma) {
using compute_t = float;
std::vector<compute_t> dgamma(H, 0.f);
std::vector<compute_t> dbeta(H, 0.f);
Expand All@@ -81,7 +85,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
dgamma[j] += y * dz;
Expand All@@ -96,7 +103,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
const compute_t dx = rsigma[i] * (dy - mdyy * y - mdy);
Expand All@@ -112,7 +122,7 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
}

template <typename InputType, typename OutputType>
void performTest(const size_t N, const size_t H) {
void performTest(const size_t N, const size_t H, const bool zero_centered_gamma) {
if (sizeof(InputType) < sizeof(OutputType)) {
GTEST_SKIP() << "LN kernel does not support OutputType > InputType";
return;
Expand DownExpand Up@@ -158,32 +168,34 @@ void performTest(const size_t N, const size_t H) {

// Forward kernel
float epsilon = 1e-5;
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto fwd_function = zero_centered_gamma ? nvte_layernorm1p_fwd : nvte_layernorm_fwd;
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Backward kernel
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto bwd_function = zero_centered_gamma ? nvte_layernorm1p_bwd : nvte_layernorm_bwd;
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
dgamma_part = Tensor(dgamma_part.shape(), dgamma_part.dtype());
dbeta_part = Tensor(dbeta_part.shape(), dbeta_part.dtype());
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Reference implementations
// use the GPU stats to tighten the tolerances
Expand All@@ -201,12 +213,13 @@ void performTest(const size_t N, const size_t H) {
rsigma.cpu_dptr<float>(),
N, H,
&ref_amax,
ref_scale);
ref_scale,
zero_centered_gamma);
compute_ref_backward(dz.cpu_dptr<WeightType>(), input.cpu_dptr<InputType>(),
mu.cpu_dptr<float>(), rsigma.cpu_dptr<float>(),
gamma.cpu_dptr<WeightType>(),
ref_dx.get(), ref_dgamma.get(), ref_dbeta.get(),
N, H);
N, H, zero_centered_gamma);

cudaDeviceSynchronize();
auto err = cudaGetLastError();
Expand DownExpand Up@@ -248,7 +261,8 @@ std::vector<std::pair<size_t, size_t>> test_cases = {{2048, 12288},

class LNTestSuite : public ::testing::TestWithParam<std::tuple<transformer_engine::DType,
transformer_engine::DType,
std::pair<size_t, size_t>>> {};
std::pair<size_t, size_t>,
bool>> {};

TEST_P(LNTestSuite, TestLN) {
using namespace transformer_engine;
Expand All@@ -257,10 +271,11 @@ TEST_P(LNTestSuite, TestLN) {
const DType input_type = std::get<0>(GetParam());
const DType output_type = std::get<1>(GetParam());
const auto size = std::get<2>(GetParam());
const bool zero_centered_gamma = std::get<3>(GetParam());

TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(input_type, InputType,
TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(output_type, OutputType,
performTest<InputType, OutputType>(size.first, size.second);
performTest<InputType, OutputType>(size.first, size.second, zero_centered_gamma);
);
);
}
Expand All@@ -271,11 +286,13 @@ INSTANTIATE_TEST_SUITE_P(
::testing::Combine(
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16),
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16, DType::kFloat8E4M3),
::testing::ValuesIn(test_cases)),
::testing::ValuesIn(test_cases),
::testing::Values(false, true)),
[](const testing::TestParamInfo<LNTestSuite::ParamType>& info) {
std::string name = test::typeName(std::get<0>(info.param)) + "X" +
test::typeName(std::get<1>(info.param)) + "X" +
std::to_string(std::get<2>(info.param).first) + "X" +
std::to_string(std::get<2>(info.param).second);
std::to_string(std::get<2>(info.param).second) + "X" +
std::to_string(std::get<3>(info.param));
return name;
});
29 changes: 21 additions & 8 deletions tests/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -434,10 +434,12 @@ def forward(self, inp, weight):
@pytest.mark.parametrize("use_fp8", [False, True])
@pytest.mark.parametrize("scale_factor", [448, 112])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm(
use_fp8: bool,
scale_factor: float,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -459,7 +461,8 @@ def forward(self, inp):
inp,
self.weight,
self.bias,
self.eps)
self.eps,
zero_centered_gamma)
return ret

class TestFP8_Layernorm(nn.Module):
Expand All@@ -482,7 +485,8 @@ def forward(self, inp):
self.eps,
self.meta,
self.fp8_tensor,
self.fp8_type)
self.fp8_type,
zero_centered_gamma)

ret = cast_from_fp8(
ret,
Expand All@@ -500,7 +504,7 @@ def forward(self, inp):
do_export(model, inp, fname, use_fp8=use_fp8)
if precision not in (torch.bfloat16, ):
# TODO: FP32 has a small threshold (1e-5)
validate_result(fname, inp, model, atol=1e-3, is_fp8=use_fp8)
validate_result(fname, inp, model, atol=4e-3, is_fp8=use_fp8)


@skip_FP8
Expand DownExpand Up@@ -646,13 +650,15 @@ def forward(self, inp):
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_linear(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -676,6 +682,7 @@ def test_export_layernorm_linear(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand All@@ -698,13 +705,15 @@ def test_export_layernorm_linear(
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_mlp(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -729,6 +738,7 @@ def test_export_layernorm_mlp(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand DownExpand Up@@ -902,14 +912,16 @@ def test_export_multihead_attention(
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
use_mask: bool,
attn_mask_type: str,
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand DownExpand Up@@ -947,7 +959,8 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling).to(device='cuda')
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
validate_result(fname, inp, model, atol=1e-3)
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Force GitHub README to respect dark mode (function() { var style = document.createElement('style'); style.textContent = ' .markdown-body { color-scheme: dark light; } .markdown-body pre { background: #161b22 !important; } .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; } .markdown-body table th, .markdown-body table td { border-color: #30363d !important; } .markdown-body img { background: #0d1117; } .markdown-body blockquote { border-left-color: #8b949e; } .markdown-body hr { border-color: #30363d; } '; document.head.appendChild(style); })(); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
79 changes: 48 additions & 31 deletions tests/cpp/operator/test_layernorm.cu
Original file line numberDiff line numberDiff line change
Expand Up@@ -49,14 +49,17 @@ template <typename InputType, typename OutputType>
void compute_ref_output(const InputType *data, const InputType *gamma, const InputType *beta,
OutputType *output, const float *mu, const float *rsigma,
const size_t N, const size_t H,
float *amax, float scale) {
float *amax, float scale, const bool zero_centered_gamma) {
using compute_t = float;
compute_t current_max = -1e100;
for (size_t i = 0 ; i < N; ++i) {
for (size_t j = 0; j < H; ++j) {
compute_t current = static_cast<compute_t>(data[i * H + j]);
compute_t tmp = (current - mu[i]) * rsigma[i] * static_cast<compute_t>(gamma[j]) +
static_cast<compute_t>(beta[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
compute_t tmp = (current - mu[i]) * rsigma[i] * g + static_cast<compute_t>(beta[j]);
output[i * H + j] = static_cast<OutputType>(tmp * scale);
current_max = fmaxf(current_max, fabsf(tmp));
}
Expand All@@ -70,7 +73,8 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
const InputType *gamma,
InputType *data_grad,
InputType *gamma_grad, InputType *beta_grad,
const size_t N, const size_t H) {
const size_t N, const size_t H,
const bool zero_centered_gamma) {
using compute_t = float;
std::vector<compute_t> dgamma(H, 0.f);
std::vector<compute_t> dbeta(H, 0.f);
Expand All@@ -81,7 +85,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
dgamma[j] += y * dz;
Expand All@@ -96,7 +103,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
const compute_t dx = rsigma[i] * (dy - mdyy * y - mdy);
Expand All@@ -112,7 +122,7 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
}

template <typename InputType, typename OutputType>
void performTest(const size_t N, const size_t H) {
void performTest(const size_t N, const size_t H, const bool zero_centered_gamma) {
if (sizeof(InputType) < sizeof(OutputType)) {
GTEST_SKIP() << "LN kernel does not support OutputType > InputType";
return;
Expand DownExpand Up@@ -158,32 +168,34 @@ void performTest(const size_t N, const size_t H) {

// Forward kernel
float epsilon = 1e-5;
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto fwd_function = zero_centered_gamma ? nvte_layernorm1p_fwd : nvte_layernorm_fwd;
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Backward kernel
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto bwd_function = zero_centered_gamma ? nvte_layernorm1p_bwd : nvte_layernorm_bwd;
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
dgamma_part = Tensor(dgamma_part.shape(), dgamma_part.dtype());
dbeta_part = Tensor(dbeta_part.shape(), dbeta_part.dtype());
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Reference implementations
// use the GPU stats to tighten the tolerances
Expand All@@ -201,12 +213,13 @@ void performTest(const size_t N, const size_t H) {
rsigma.cpu_dptr<float>(),
N, H,
&ref_amax,
ref_scale);
ref_scale,
zero_centered_gamma);
compute_ref_backward(dz.cpu_dptr<WeightType>(), input.cpu_dptr<InputType>(),
mu.cpu_dptr<float>(), rsigma.cpu_dptr<float>(),
gamma.cpu_dptr<WeightType>(),
ref_dx.get(), ref_dgamma.get(), ref_dbeta.get(),
N, H);
N, H, zero_centered_gamma);

cudaDeviceSynchronize();
auto err = cudaGetLastError();
Expand DownExpand Up@@ -248,7 +261,8 @@ std::vector<std::pair<size_t, size_t>> test_cases = {{2048, 12288},

class LNTestSuite : public ::testing::TestWithParam<std::tuple<transformer_engine::DType,
transformer_engine::DType,
std::pair<size_t, size_t>>> {};
std::pair<size_t, size_t>,
bool>> {};

TEST_P(LNTestSuite, TestLN) {
using namespace transformer_engine;
Expand All@@ -257,10 +271,11 @@ TEST_P(LNTestSuite, TestLN) {
const DType input_type = std::get<0>(GetParam());
const DType output_type = std::get<1>(GetParam());
const auto size = std::get<2>(GetParam());
const bool zero_centered_gamma = std::get<3>(GetParam());

TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(input_type, InputType,
TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(output_type, OutputType,
performTest<InputType, OutputType>(size.first, size.second);
performTest<InputType, OutputType>(size.first, size.second, zero_centered_gamma);
);
);
}
Expand All@@ -271,11 +286,13 @@ INSTANTIATE_TEST_SUITE_P(
::testing::Combine(
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16),
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16, DType::kFloat8E4M3),
::testing::ValuesIn(test_cases)),
::testing::ValuesIn(test_cases),
::testing::Values(false, true)),
[](const testing::TestParamInfo<LNTestSuite::ParamType>& info) {
std::string name = test::typeName(std::get<0>(info.param)) + "X" +
test::typeName(std::get<1>(info.param)) + "X" +
std::to_string(std::get<2>(info.param).first) + "X" +
std::to_string(std::get<2>(info.param).second);
std::to_string(std::get<2>(info.param).second) + "X" +
std::to_string(std::get<3>(info.param));
return name;
});
29 changes: 21 additions & 8 deletions tests/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -434,10 +434,12 @@ def forward(self, inp, weight):
@pytest.mark.parametrize("use_fp8", [False, True])
@pytest.mark.parametrize("scale_factor", [448, 112])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm(
use_fp8: bool,
scale_factor: float,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -459,7 +461,8 @@ def forward(self, inp):
inp,
self.weight,
self.bias,
self.eps)
self.eps,
zero_centered_gamma)
return ret

class TestFP8_Layernorm(nn.Module):
Expand All@@ -482,7 +485,8 @@ def forward(self, inp):
self.eps,
self.meta,
self.fp8_tensor,
self.fp8_type)
self.fp8_type,
zero_centered_gamma)

ret = cast_from_fp8(
ret,
Expand All@@ -500,7 +504,7 @@ def forward(self, inp):
do_export(model, inp, fname, use_fp8=use_fp8)
if precision not in (torch.bfloat16, ):
# TODO: FP32 has a small threshold (1e-5)
validate_result(fname, inp, model, atol=1e-3, is_fp8=use_fp8)
validate_result(fname, inp, model, atol=4e-3, is_fp8=use_fp8)


@skip_FP8
Expand DownExpand Up@@ -646,13 +650,15 @@ def forward(self, inp):
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_linear(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -676,6 +682,7 @@ def test_export_layernorm_linear(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand All@@ -698,13 +705,15 @@ def test_export_layernorm_linear(
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_mlp(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -729,6 +738,7 @@ def test_export_layernorm_mlp(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand DownExpand Up@@ -902,14 +912,16 @@ def test_export_multihead_attention(
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
use_mask: bool,
attn_mask_type: str,
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand DownExpand Up@@ -947,7 +959,8 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling).to(device='cuda')
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
validate_result(fname, inp, model, atol=1e-3)
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Highlight search terms from Google/DuckDuckGo/Bing referrer (function() { var ref = document.referrer; var terms = []; if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) { var url = new URL(ref); var q = url.searchParams.get('q') || url.searchParams.get('p'); if (q) { terms = q.split(/\s+/).filter(function(t) { return t.length > 2; }); } } if (terms.length === 0) return; var style = document.createElement('style'); style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }'; document.head.appendChild(style); function highlight(node) { if (node.nodeType === 3) { // text node var text = node.textContent; var found = false; terms.forEach(function(term) { var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\]\\]/g, '\\') + ')', 'gi'); if (regex.test(text)) { found = true; var frag = document.createDocumentFragment(); var parts = text.split(regex); parts.forEach(function(part, i) { if (i % 2 === 0) { frag.appendChild(document.createTextNode(part)); } else { var span = document.createElement('span'); span.className = 'userscript-highlight'; span.textContent = part; frag.appendChild(span); } }); node.parentNode.replaceChild(frag, node); } }); } else if (node.nodeType === 1 && node.childNodes) { // element var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT']; if (!skipTags.includes(node.tagName)) { Array.from(node.childNodes).forEach(highlight); } } } highlight(document.body); // Re-highlight on dynamic content var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1 || node.nodeType === 3) highlight(node); }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
79 changes: 48 additions & 31 deletions tests/cpp/operator/test_layernorm.cu
Original file line numberDiff line numberDiff line change
Expand Up@@ -49,14 +49,17 @@ template <typename InputType, typename OutputType>
void compute_ref_output(const InputType *data, const InputType *gamma, const InputType *beta,
OutputType *output, const float *mu, const float *rsigma,
const size_t N, const size_t H,
float *amax, float scale) {
float *amax, float scale, const bool zero_centered_gamma) {
using compute_t = float;
compute_t current_max = -1e100;
for (size_t i = 0 ; i < N; ++i) {
for (size_t j = 0; j < H; ++j) {
compute_t current = static_cast<compute_t>(data[i * H + j]);
compute_t tmp = (current - mu[i]) * rsigma[i] * static_cast<compute_t>(gamma[j]) +
static_cast<compute_t>(beta[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
compute_t tmp = (current - mu[i]) * rsigma[i] * g + static_cast<compute_t>(beta[j]);
output[i * H + j] = static_cast<OutputType>(tmp * scale);
current_max = fmaxf(current_max, fabsf(tmp));
}
Expand All@@ -70,7 +73,8 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
const InputType *gamma,
InputType *data_grad,
InputType *gamma_grad, InputType *beta_grad,
const size_t N, const size_t H) {
const size_t N, const size_t H,
const bool zero_centered_gamma) {
using compute_t = float;
std::vector<compute_t> dgamma(H, 0.f);
std::vector<compute_t> dbeta(H, 0.f);
Expand All@@ -81,7 +85,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
dgamma[j] += y * dz;
Expand All@@ -96,7 +103,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
const compute_t dx = rsigma[i] * (dy - mdyy * y - mdy);
Expand All@@ -112,7 +122,7 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
}

template <typename InputType, typename OutputType>
void performTest(const size_t N, const size_t H) {
void performTest(const size_t N, const size_t H, const bool zero_centered_gamma) {
if (sizeof(InputType) < sizeof(OutputType)) {
GTEST_SKIP() << "LN kernel does not support OutputType > InputType";
return;
Expand DownExpand Up@@ -158,32 +168,34 @@ void performTest(const size_t N, const size_t H) {

// Forward kernel
float epsilon = 1e-5;
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto fwd_function = zero_centered_gamma ? nvte_layernorm1p_fwd : nvte_layernorm_fwd;
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Backward kernel
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto bwd_function = zero_centered_gamma ? nvte_layernorm1p_bwd : nvte_layernorm_bwd;
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
dgamma_part = Tensor(dgamma_part.shape(), dgamma_part.dtype());
dbeta_part = Tensor(dbeta_part.shape(), dbeta_part.dtype());
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Reference implementations
// use the GPU stats to tighten the tolerances
Expand All@@ -201,12 +213,13 @@ void performTest(const size_t N, const size_t H) {
rsigma.cpu_dptr<float>(),
N, H,
&ref_amax,
ref_scale);
ref_scale,
zero_centered_gamma);
compute_ref_backward(dz.cpu_dptr<WeightType>(), input.cpu_dptr<InputType>(),
mu.cpu_dptr<float>(), rsigma.cpu_dptr<float>(),
gamma.cpu_dptr<WeightType>(),
ref_dx.get(), ref_dgamma.get(), ref_dbeta.get(),
N, H);
N, H, zero_centered_gamma);

cudaDeviceSynchronize();
auto err = cudaGetLastError();
Expand DownExpand Up@@ -248,7 +261,8 @@ std::vector<std::pair<size_t, size_t>> test_cases = {{2048, 12288},

class LNTestSuite : public ::testing::TestWithParam<std::tuple<transformer_engine::DType,
transformer_engine::DType,
std::pair<size_t, size_t>>> {};
std::pair<size_t, size_t>,
bool>> {};

TEST_P(LNTestSuite, TestLN) {
using namespace transformer_engine;
Expand All@@ -257,10 +271,11 @@ TEST_P(LNTestSuite, TestLN) {
const DType input_type = std::get<0>(GetParam());
const DType output_type = std::get<1>(GetParam());
const auto size = std::get<2>(GetParam());
const bool zero_centered_gamma = std::get<3>(GetParam());

TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(input_type, InputType,
TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(output_type, OutputType,
performTest<InputType, OutputType>(size.first, size.second);
performTest<InputType, OutputType>(size.first, size.second, zero_centered_gamma);
);
);
}
Expand All@@ -271,11 +286,13 @@ INSTANTIATE_TEST_SUITE_P(
::testing::Combine(
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16),
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16, DType::kFloat8E4M3),
::testing::ValuesIn(test_cases)),
::testing::ValuesIn(test_cases),
::testing::Values(false, true)),
[](const testing::TestParamInfo<LNTestSuite::ParamType>& info) {
std::string name = test::typeName(std::get<0>(info.param)) + "X" +
test::typeName(std::get<1>(info.param)) + "X" +
std::to_string(std::get<2>(info.param).first) + "X" +
std::to_string(std::get<2>(info.param).second);
std::to_string(std::get<2>(info.param).second) + "X" +
std::to_string(std::get<3>(info.param));
return name;
});
29 changes: 21 additions & 8 deletions tests/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -434,10 +434,12 @@ def forward(self, inp, weight):
@pytest.mark.parametrize("use_fp8", [False, True])
@pytest.mark.parametrize("scale_factor", [448, 112])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm(
use_fp8: bool,
scale_factor: float,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -459,7 +461,8 @@ def forward(self, inp):
inp,
self.weight,
self.bias,
self.eps)
self.eps,
zero_centered_gamma)
return ret

class TestFP8_Layernorm(nn.Module):
Expand All@@ -482,7 +485,8 @@ def forward(self, inp):
self.eps,
self.meta,
self.fp8_tensor,
self.fp8_type)
self.fp8_type,
zero_centered_gamma)

ret = cast_from_fp8(
ret,
Expand All@@ -500,7 +504,7 @@ def forward(self, inp):
do_export(model, inp, fname, use_fp8=use_fp8)
if precision not in (torch.bfloat16, ):
# TODO: FP32 has a small threshold (1e-5)
validate_result(fname, inp, model, atol=1e-3, is_fp8=use_fp8)
validate_result(fname, inp, model, atol=4e-3, is_fp8=use_fp8)


@skip_FP8
Expand DownExpand Up@@ -646,13 +650,15 @@ def forward(self, inp):
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_linear(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -676,6 +682,7 @@ def test_export_layernorm_linear(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand All@@ -698,13 +705,15 @@ def test_export_layernorm_linear(
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_mlp(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -729,6 +738,7 @@ def test_export_layernorm_mlp(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand DownExpand Up@@ -902,14 +912,16 @@ def test_export_multihead_attention(
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
use_mask: bool,
attn_mask_type: str,
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand DownExpand Up@@ -947,7 +959,8 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling).to(device='cuda')
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
validate_result(fname, inp, model, atol=1e-3)
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Strip utm_, fbclid, gclid, etc. from all links on page (function() { var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content', 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid', 'ref', 'ref_src', 'source', 'medium', 'campaign']; function cleanUrl(url) { try { var u = new URL(url, window.location.origin); var changed = false; trackingParams.forEach(function(p) { if (u.searchParams.has(p)) { u.searchParams.delete(p); changed = true; } }); return changed ? u.toString() : url; } catch (e) { return url; } } function cleanLinks() { document.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } cleanLinks(); var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1) { if (node.tagName === 'A') cleanLinks(); node.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
Skip to content
Merged
79 changes: 48 additions & 31 deletions tests/cpp/operator/test_layernorm.cu
Original file line numberDiff line numberDiff line change
Expand Up@@ -49,14 +49,17 @@ template <typename InputType, typename OutputType>
void compute_ref_output(const InputType *data, const InputType *gamma, const InputType *beta,
OutputType *output, const float *mu, const float *rsigma,
const size_t N, const size_t H,
float *amax, float scale) {
float *amax, float scale, const bool zero_centered_gamma) {
using compute_t = float;
compute_t current_max = -1e100;
for (size_t i = 0 ; i < N; ++i) {
for (size_t j = 0; j < H; ++j) {
compute_t current = static_cast<compute_t>(data[i * H + j]);
compute_t tmp = (current - mu[i]) * rsigma[i] * static_cast<compute_t>(gamma[j]) +
static_cast<compute_t>(beta[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
compute_t tmp = (current - mu[i]) * rsigma[i] * g + static_cast<compute_t>(beta[j]);
output[i * H + j] = static_cast<OutputType>(tmp * scale);
current_max = fmaxf(current_max, fabsf(tmp));
}
Expand All@@ -70,7 +73,8 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
const InputType *gamma,
InputType *data_grad,
InputType *gamma_grad, InputType *beta_grad,
const size_t N, const size_t H) {
const size_t N, const size_t H,
const bool zero_centered_gamma) {
using compute_t = float;
std::vector<compute_t> dgamma(H, 0.f);
std::vector<compute_t> dbeta(H, 0.f);
Expand All@@ -81,7 +85,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
dgamma[j] += y * dz;
Expand All@@ -96,7 +103,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
const compute_t dx = rsigma[i] * (dy - mdyy * y - mdy);
Expand All@@ -112,7 +122,7 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
}

template <typename InputType, typename OutputType>
void performTest(const size_t N, const size_t H) {
void performTest(const size_t N, const size_t H, const bool zero_centered_gamma) {
if (sizeof(InputType) < sizeof(OutputType)) {
GTEST_SKIP() << "LN kernel does not support OutputType > InputType";
return;
Expand DownExpand Up@@ -158,32 +168,34 @@ void performTest(const size_t N, const size_t H) {

// Forward kernel
float epsilon = 1e-5;
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto fwd_function = zero_centered_gamma ? nvte_layernorm1p_fwd : nvte_layernorm_fwd;
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Backward kernel
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto bwd_function = zero_centered_gamma ? nvte_layernorm1p_bwd : nvte_layernorm_bwd;
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
dgamma_part = Tensor(dgamma_part.shape(), dgamma_part.dtype());
dbeta_part = Tensor(dbeta_part.shape(), dbeta_part.dtype());
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Reference implementations
// use the GPU stats to tighten the tolerances
Expand All@@ -201,12 +213,13 @@ void performTest(const size_t N, const size_t H) {
rsigma.cpu_dptr<float>(),
N, H,
&ref_amax,
ref_scale);
ref_scale,
zero_centered_gamma);
compute_ref_backward(dz.cpu_dptr<WeightType>(), input.cpu_dptr<InputType>(),
mu.cpu_dptr<float>(), rsigma.cpu_dptr<float>(),
gamma.cpu_dptr<WeightType>(),
ref_dx.get(), ref_dgamma.get(), ref_dbeta.get(),
N, H);
N, H, zero_centered_gamma);

cudaDeviceSynchronize();
auto err = cudaGetLastError();
Expand DownExpand Up@@ -248,7 +261,8 @@ std::vector<std::pair<size_t, size_t>> test_cases = {{2048, 12288},

class LNTestSuite : public ::testing::TestWithParam<std::tuple<transformer_engine::DType,
transformer_engine::DType,
std::pair<size_t, size_t>>> {};
std::pair<size_t, size_t>,
bool>> {};

TEST_P(LNTestSuite, TestLN) {
using namespace transformer_engine;
Expand All@@ -257,10 +271,11 @@ TEST_P(LNTestSuite, TestLN) {
const DType input_type = std::get<0>(GetParam());
const DType output_type = std::get<1>(GetParam());
const auto size = std::get<2>(GetParam());
const bool zero_centered_gamma = std::get<3>(GetParam());

TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(input_type, InputType,
TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(output_type, OutputType,
performTest<InputType, OutputType>(size.first, size.second);
performTest<InputType, OutputType>(size.first, size.second, zero_centered_gamma);
);
);
}
Expand All@@ -271,11 +286,13 @@ INSTANTIATE_TEST_SUITE_P(
::testing::Combine(
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16),
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16, DType::kFloat8E4M3),
::testing::ValuesIn(test_cases)),
::testing::ValuesIn(test_cases),
::testing::Values(false, true)),
[](const testing::TestParamInfo<LNTestSuite::ParamType>& info) {
std::string name = test::typeName(std::get<0>(info.param)) + "X" +
test::typeName(std::get<1>(info.param)) + "X" +
std::to_string(std::get<2>(info.param).first) + "X" +
std::to_string(std::get<2>(info.param).second);
std::to_string(std::get<2>(info.param).second) + "X" +
std::to_string(std::get<3>(info.param));
return name;
});
29 changes: 21 additions & 8 deletions tests/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -434,10 +434,12 @@ def forward(self, inp, weight):
@pytest.mark.parametrize("use_fp8", [False, True])
@pytest.mark.parametrize("scale_factor", [448, 112])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm(
use_fp8: bool,
scale_factor: float,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -459,7 +461,8 @@ def forward(self, inp):
inp,
self.weight,
self.bias,
self.eps)
self.eps,
zero_centered_gamma)
return ret

class TestFP8_Layernorm(nn.Module):
Expand All@@ -482,7 +485,8 @@ def forward(self, inp):
self.eps,
self.meta,
self.fp8_tensor,
self.fp8_type)
self.fp8_type,
zero_centered_gamma)

ret = cast_from_fp8(
ret,
Expand All@@ -500,7 +504,7 @@ def forward(self, inp):
do_export(model, inp, fname, use_fp8=use_fp8)
if precision not in (torch.bfloat16, ):
# TODO: FP32 has a small threshold (1e-5)
validate_result(fname, inp, model, atol=1e-3, is_fp8=use_fp8)
validate_result(fname, inp, model, atol=4e-3, is_fp8=use_fp8)


@skip_FP8
Expand DownExpand Up@@ -646,13 +650,15 @@ def forward(self, inp):
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_linear(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -676,6 +682,7 @@ def test_export_layernorm_linear(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand All@@ -698,13 +705,15 @@ def test_export_layernorm_linear(
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_mlp(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -729,6 +738,7 @@ def test_export_layernorm_mlp(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand DownExpand Up@@ -902,14 +912,16 @@ def test_export_multihead_attention(
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
use_mask: bool,
attn_mask_type: str,
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand DownExpand Up@@ -947,7 +959,8 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling).to(device='cuda')
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
validate_result(fname, inp, model, atol=1e-3)
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Auto-enable theater mode on YouTube (function() { function tryTheater() { var btn = document.querySelector('button[aria-label="Theater mode"], ytd-player #player button[title="Theater mode"]'); if (btn && !btn.classList.contains('activated')) { btn.click(); } } // Try immediately tryTheater(); // Try after navigation (SPA) var lastUrl = location.href; setInterval(function() { if (location.href !== lastUrl) { lastUrl = location.href; setTimeout(tryTheater, 500); } }, 1000); // Also try on player load var observer = new MutationObserver(tryTheater); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
79 changes: 48 additions & 31 deletions tests/cpp/operator/test_layernorm.cu
Original file line numberDiff line numberDiff line change
Expand Up@@ -49,14 +49,17 @@ template <typename InputType, typename OutputType>
void compute_ref_output(const InputType *data, const InputType *gamma, const InputType *beta,
OutputType *output, const float *mu, const float *rsigma,
const size_t N, const size_t H,
float *amax, float scale) {
float *amax, float scale, const bool zero_centered_gamma) {
using compute_t = float;
compute_t current_max = -1e100;
for (size_t i = 0 ; i < N; ++i) {
for (size_t j = 0; j < H; ++j) {
compute_t current = static_cast<compute_t>(data[i * H + j]);
compute_t tmp = (current - mu[i]) * rsigma[i] * static_cast<compute_t>(gamma[j]) +
static_cast<compute_t>(beta[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
compute_t tmp = (current - mu[i]) * rsigma[i] * g + static_cast<compute_t>(beta[j]);
output[i * H + j] = static_cast<OutputType>(tmp * scale);
current_max = fmaxf(current_max, fabsf(tmp));
}
Expand All@@ -70,7 +73,8 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
const InputType *gamma,
InputType *data_grad,
InputType *gamma_grad, InputType *beta_grad,
const size_t N, const size_t H) {
const size_t N, const size_t H,
const bool zero_centered_gamma) {
using compute_t = float;
std::vector<compute_t> dgamma(H, 0.f);
std::vector<compute_t> dbeta(H, 0.f);
Expand All@@ -81,7 +85,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
dgamma[j] += y * dz;
Expand All@@ -96,7 +103,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
const compute_t dx = rsigma[i] * (dy - mdyy * y - mdy);
Expand All@@ -112,7 +122,7 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
}

template <typename InputType, typename OutputType>
void performTest(const size_t N, const size_t H) {
void performTest(const size_t N, const size_t H, const bool zero_centered_gamma) {
if (sizeof(InputType) < sizeof(OutputType)) {
GTEST_SKIP() << "LN kernel does not support OutputType > InputType";
return;
Expand DownExpand Up@@ -158,32 +168,34 @@ void performTest(const size_t N, const size_t H) {

// Forward kernel
float epsilon = 1e-5;
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto fwd_function = zero_centered_gamma ? nvte_layernorm1p_fwd : nvte_layernorm_fwd;
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Backward kernel
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto bwd_function = zero_centered_gamma ? nvte_layernorm1p_bwd : nvte_layernorm_bwd;
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
dgamma_part = Tensor(dgamma_part.shape(), dgamma_part.dtype());
dbeta_part = Tensor(dbeta_part.shape(), dbeta_part.dtype());
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Reference implementations
// use the GPU stats to tighten the tolerances
Expand All@@ -201,12 +213,13 @@ void performTest(const size_t N, const size_t H) {
rsigma.cpu_dptr<float>(),
N, H,
&ref_amax,
ref_scale);
ref_scale,
zero_centered_gamma);
compute_ref_backward(dz.cpu_dptr<WeightType>(), input.cpu_dptr<InputType>(),
mu.cpu_dptr<float>(), rsigma.cpu_dptr<float>(),
gamma.cpu_dptr<WeightType>(),
ref_dx.get(), ref_dgamma.get(), ref_dbeta.get(),
N, H);
N, H, zero_centered_gamma);

cudaDeviceSynchronize();
auto err = cudaGetLastError();
Expand DownExpand Up@@ -248,7 +261,8 @@ std::vector<std::pair<size_t, size_t>> test_cases = {{2048, 12288},

class LNTestSuite : public ::testing::TestWithParam<std::tuple<transformer_engine::DType,
transformer_engine::DType,
std::pair<size_t, size_t>>> {};
std::pair<size_t, size_t>,
bool>> {};

TEST_P(LNTestSuite, TestLN) {
using namespace transformer_engine;
Expand All@@ -257,10 +271,11 @@ TEST_P(LNTestSuite, TestLN) {
const DType input_type = std::get<0>(GetParam());
const DType output_type = std::get<1>(GetParam());
const auto size = std::get<2>(GetParam());
const bool zero_centered_gamma = std::get<3>(GetParam());

TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(input_type, InputType,
TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(output_type, OutputType,
performTest<InputType, OutputType>(size.first, size.second);
performTest<InputType, OutputType>(size.first, size.second, zero_centered_gamma);
);
);
}
Expand All@@ -271,11 +286,13 @@ INSTANTIATE_TEST_SUITE_P(
::testing::Combine(
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16),
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16, DType::kFloat8E4M3),
::testing::ValuesIn(test_cases)),
::testing::ValuesIn(test_cases),
::testing::Values(false, true)),
[](const testing::TestParamInfo<LNTestSuite::ParamType>& info) {
std::string name = test::typeName(std::get<0>(info.param)) + "X" +
test::typeName(std::get<1>(info.param)) + "X" +
std::to_string(std::get<2>(info.param).first) + "X" +
std::to_string(std::get<2>(info.param).second);
std::to_string(std::get<2>(info.param).second) + "X" +
std::to_string(std::get<3>(info.param));
return name;
});
29 changes: 21 additions & 8 deletions tests/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -434,10 +434,12 @@ def forward(self, inp, weight):
@pytest.mark.parametrize("use_fp8", [False, True])
@pytest.mark.parametrize("scale_factor", [448, 112])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm(
use_fp8: bool,
scale_factor: float,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -459,7 +461,8 @@ def forward(self, inp):
inp,
self.weight,
self.bias,
self.eps)
self.eps,
zero_centered_gamma)
return ret

class TestFP8_Layernorm(nn.Module):
Expand All@@ -482,7 +485,8 @@ def forward(self, inp):
self.eps,
self.meta,
self.fp8_tensor,
self.fp8_type)
self.fp8_type,
zero_centered_gamma)

ret = cast_from_fp8(
ret,
Expand All@@ -500,7 +504,7 @@ def forward(self, inp):
do_export(model, inp, fname, use_fp8=use_fp8)
if precision not in (torch.bfloat16, ):
# TODO: FP32 has a small threshold (1e-5)
validate_result(fname, inp, model, atol=1e-3, is_fp8=use_fp8)
validate_result(fname, inp, model, atol=4e-3, is_fp8=use_fp8)


@skip_FP8
Expand DownExpand Up@@ -646,13 +650,15 @@ def forward(self, inp):
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_linear(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -676,6 +682,7 @@ def test_export_layernorm_linear(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand All@@ -698,13 +705,15 @@ def test_export_layernorm_linear(
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_mlp(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -729,6 +738,7 @@ def test_export_layernorm_mlp(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand DownExpand Up@@ -902,14 +912,16 @@ def test_export_multihead_attention(
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
use_mask: bool,
attn_mask_type: str,
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand DownExpand Up@@ -947,7 +959,8 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling).to(device='cuda')
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
validate_result(fname, inp, model, atol=1e-3)
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Remove or un-stick sticky/fixed headers that block content (function() { function unstick() { document.querySelectorAll('header, nav, [role="banner"], .header, .navbar, .sticky, .fixed-top, [style*="position: fixed"], [style*="position:sticky"]').forEach(function(el) { if (el.style.position === 'fixed' || el.style.position === 'sticky' || getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') { el.style.position = 'static'; el.style.top = 'auto'; el.style.zIndex = 'auto'; } }); } unstick(); var observer = new MutationObserver(unstick); observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] }); })(); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
79 changes: 48 additions & 31 deletions tests/cpp/operator/test_layernorm.cu
Original file line numberDiff line numberDiff line change
Expand Up@@ -49,14 +49,17 @@ template <typename InputType, typename OutputType>
void compute_ref_output(const InputType *data, const InputType *gamma, const InputType *beta,
OutputType *output, const float *mu, const float *rsigma,
const size_t N, const size_t H,
float *amax, float scale) {
float *amax, float scale, const bool zero_centered_gamma) {
using compute_t = float;
compute_t current_max = -1e100;
for (size_t i = 0 ; i < N; ++i) {
for (size_t j = 0; j < H; ++j) {
compute_t current = static_cast<compute_t>(data[i * H + j]);
compute_t tmp = (current - mu[i]) * rsigma[i] * static_cast<compute_t>(gamma[j]) +
static_cast<compute_t>(beta[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
compute_t tmp = (current - mu[i]) * rsigma[i] * g + static_cast<compute_t>(beta[j]);
output[i * H + j] = static_cast<OutputType>(tmp * scale);
current_max = fmaxf(current_max, fabsf(tmp));
}
Expand All@@ -70,7 +73,8 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
const InputType *gamma,
InputType *data_grad,
InputType *gamma_grad, InputType *beta_grad,
const size_t N, const size_t H) {
const size_t N, const size_t H,
const bool zero_centered_gamma) {
using compute_t = float;
std::vector<compute_t> dgamma(H, 0.f);
std::vector<compute_t> dbeta(H, 0.f);
Expand All@@ -81,7 +85,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
dgamma[j] += y * dz;
Expand All@@ -96,7 +103,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
const compute_t dx = rsigma[i] * (dy - mdyy * y - mdy);
Expand All@@ -112,7 +122,7 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
}

template <typename InputType, typename OutputType>
void performTest(const size_t N, const size_t H) {
void performTest(const size_t N, const size_t H, const bool zero_centered_gamma) {
if (sizeof(InputType) < sizeof(OutputType)) {
GTEST_SKIP() << "LN kernel does not support OutputType > InputType";
return;
Expand DownExpand Up@@ -158,32 +168,34 @@ void performTest(const size_t N, const size_t H) {

// Forward kernel
float epsilon = 1e-5;
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto fwd_function = zero_centered_gamma ? nvte_layernorm1p_fwd : nvte_layernorm_fwd;
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Backward kernel
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto bwd_function = zero_centered_gamma ? nvte_layernorm1p_bwd : nvte_layernorm_bwd;
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
dgamma_part = Tensor(dgamma_part.shape(), dgamma_part.dtype());
dbeta_part = Tensor(dbeta_part.shape(), dbeta_part.dtype());
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Reference implementations
// use the GPU stats to tighten the tolerances
Expand All@@ -201,12 +213,13 @@ void performTest(const size_t N, const size_t H) {
rsigma.cpu_dptr<float>(),
N, H,
&ref_amax,
ref_scale);
ref_scale,
zero_centered_gamma);
compute_ref_backward(dz.cpu_dptr<WeightType>(), input.cpu_dptr<InputType>(),
mu.cpu_dptr<float>(), rsigma.cpu_dptr<float>(),
gamma.cpu_dptr<WeightType>(),
ref_dx.get(), ref_dgamma.get(), ref_dbeta.get(),
N, H);
N, H, zero_centered_gamma);

cudaDeviceSynchronize();
auto err = cudaGetLastError();
Expand DownExpand Up@@ -248,7 +261,8 @@ std::vector<std::pair<size_t, size_t>> test_cases = {{2048, 12288},

class LNTestSuite : public ::testing::TestWithParam<std::tuple<transformer_engine::DType,
transformer_engine::DType,
std::pair<size_t, size_t>>> {};
std::pair<size_t, size_t>,
bool>> {};

TEST_P(LNTestSuite, TestLN) {
using namespace transformer_engine;
Expand All@@ -257,10 +271,11 @@ TEST_P(LNTestSuite, TestLN) {
const DType input_type = std::get<0>(GetParam());
const DType output_type = std::get<1>(GetParam());
const auto size = std::get<2>(GetParam());
const bool zero_centered_gamma = std::get<3>(GetParam());

TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(input_type, InputType,
TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(output_type, OutputType,
performTest<InputType, OutputType>(size.first, size.second);
performTest<InputType, OutputType>(size.first, size.second, zero_centered_gamma);
);
);
}
Expand All@@ -271,11 +286,13 @@ INSTANTIATE_TEST_SUITE_P(
::testing::Combine(
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16),
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16, DType::kFloat8E4M3),
::testing::ValuesIn(test_cases)),
::testing::ValuesIn(test_cases),
::testing::Values(false, true)),
[](const testing::TestParamInfo<LNTestSuite::ParamType>& info) {
std::string name = test::typeName(std::get<0>(info.param)) + "X" +
test::typeName(std::get<1>(info.param)) + "X" +
std::to_string(std::get<2>(info.param).first) + "X" +
std::to_string(std::get<2>(info.param).second);
std::to_string(std::get<2>(info.param).second) + "X" +
std::to_string(std::get<3>(info.param));
return name;
});
29 changes: 21 additions & 8 deletions tests/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -434,10 +434,12 @@ def forward(self, inp, weight):
@pytest.mark.parametrize("use_fp8", [False, True])
@pytest.mark.parametrize("scale_factor", [448, 112])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm(
use_fp8: bool,
scale_factor: float,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -459,7 +461,8 @@ def forward(self, inp):
inp,
self.weight,
self.bias,
self.eps)
self.eps,
zero_centered_gamma)
return ret

class TestFP8_Layernorm(nn.Module):
Expand All@@ -482,7 +485,8 @@ def forward(self, inp):
self.eps,
self.meta,
self.fp8_tensor,
self.fp8_type)
self.fp8_type,
zero_centered_gamma)

ret = cast_from_fp8(
ret,
Expand All@@ -500,7 +504,7 @@ def forward(self, inp):
do_export(model, inp, fname, use_fp8=use_fp8)
if precision not in (torch.bfloat16, ):
# TODO: FP32 has a small threshold (1e-5)
validate_result(fname, inp, model, atol=1e-3, is_fp8=use_fp8)
validate_result(fname, inp, model, atol=4e-3, is_fp8=use_fp8)


@skip_FP8
Expand DownExpand Up@@ -646,13 +650,15 @@ def forward(self, inp):
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_linear(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -676,6 +682,7 @@ def test_export_layernorm_linear(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand All@@ -698,13 +705,15 @@ def test_export_layernorm_linear(
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_mlp(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -729,6 +738,7 @@ def test_export_layernorm_mlp(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand DownExpand Up@@ -902,14 +912,16 @@ def test_export_multihead_attention(
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
use_mask: bool,
attn_mask_type: str,
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand DownExpand Up@@ -947,7 +959,8 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling).to(device='cuda')
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
validate_result(fname, inp, model, atol=1e-3)
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Universal Dark Mode - works on any site (function() { var enabled = true; function applyDarkMode() { if (!enabled) return; // Create style element if it doesn't exist var style = document.getElementById('universal-dark-mode-style'); if (!style) { style = document.createElement('style'); style.id = 'universal-dark-mode-style'; document.head.appendChild(style); } // Dark mode CSS - inverts colors but preserves images/video style.textContent = ' /* Invert everything except media */ html { filter: invert(1) hue-rotate(180deg) !important; background: #1a1a2e !important; } /* Restore images, videos, iframes, canvas */ img, video, iframe, canvas, svg, picture, [style*="background-image"] { filter: invert(1) hue-rotate(180deg) !important; } /* Preserve specific elements that should not be inverted */ .no-dark-mode, .no-dark-mode *, [data-theme="light"], [data-theme="light"], .ace_editor, .ace_editor *, .CodeMirror, .CodeMirror *, .monaco-editor, .monaco-editor *, .markdown-body pre, .markdown-body pre *, .highlight, .highlight *, pre code, pre code * { filter: none !important; } /* Fix common UI elements */ .modal, .popup, .dropdown-menu, .tooltip, .popover { filter: invert(1) hue-rotate(180deg) !important; background: #2d2d44 !important; border-color: #444 !important; } /* Scrollbars */ ::-webkit-scrollbar { background: #1a1a2e !important; } ::-webkit-scrollbar-thumb { background: #444 !important; } ::-webkit-scrollbar-thumb:hover { background: #555 !important; } /* Selection */ ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; } ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; } '; } function removeDarkMode() { var style = document.getElementById('universal-dark-mode-style'); if (style) style.remove(); } // Toggle with Alt+Shift+D document.addEventListener('keydown', function(e) { if (e.altKey && e.shiftKey && e.key === 'D') { e.preventDefault(); enabled = !enabled; if (enabled) { applyDarkMode(); console.log('[Universal Dark Mode] Enabled'); } else { removeDarkMode(); console.log('[Universal Dark Mode] Disabled'); } } }); // Apply on load applyDarkMode(); // Re-apply on dynamic content var observer = new MutationObserver(function(mutations) { if (enabled && !document.getElementById('universal-dark-mode-style')) { applyDarkMode(); } }); observer.observe(document.head, { childList: true }); console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle'); })(); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
Skip to content
Merged
79 changes: 48 additions & 31 deletions tests/cpp/operator/test_layernorm.cu
Original file line numberDiff line numberDiff line change
Expand Up@@ -49,14 +49,17 @@ template <typename InputType, typename OutputType>
void compute_ref_output(const InputType *data, const InputType *gamma, const InputType *beta,
OutputType *output, const float *mu, const float *rsigma,
const size_t N, const size_t H,
float *amax, float scale) {
float *amax, float scale, const bool zero_centered_gamma) {
using compute_t = float;
compute_t current_max = -1e100;
for (size_t i = 0 ; i < N; ++i) {
for (size_t j = 0; j < H; ++j) {
compute_t current = static_cast<compute_t>(data[i * H + j]);
compute_t tmp = (current - mu[i]) * rsigma[i] * static_cast<compute_t>(gamma[j]) +
static_cast<compute_t>(beta[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
compute_t tmp = (current - mu[i]) * rsigma[i] * g + static_cast<compute_t>(beta[j]);
output[i * H + j] = static_cast<OutputType>(tmp * scale);
current_max = fmaxf(current_max, fabsf(tmp));
}
Expand All@@ -70,7 +73,8 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
const InputType *gamma,
InputType *data_grad,
InputType *gamma_grad, InputType *beta_grad,
const size_t N, const size_t H) {
const size_t N, const size_t H,
const bool zero_centered_gamma) {
using compute_t = float;
std::vector<compute_t> dgamma(H, 0.f);
std::vector<compute_t> dbeta(H, 0.f);
Expand All@@ -81,7 +85,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
dgamma[j] += y * dz;
Expand All@@ -96,7 +103,10 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
for (size_t j = 0; j < H; ++j) {
const compute_t x = static_cast<compute_t>(data[i * H + j]);
const compute_t y = (x - mu[i]) * rsigma[i];
const compute_t g = static_cast<compute_t>(gamma[j]);
compute_t g = static_cast<compute_t>(gamma[j]);
if (zero_centered_gamma) {
g += 1;
}
const compute_t dz = static_cast<compute_t>(output_grad[i * H + j]);
const compute_t dy = g * dz;
const compute_t dx = rsigma[i] * (dy - mdyy * y - mdy);
Expand All@@ -112,7 +122,7 @@ void compute_ref_backward(const OutputType *output_grad, const InputType *data,
}

template <typename InputType, typename OutputType>
void performTest(const size_t N, const size_t H) {
void performTest(const size_t N, const size_t H, const bool zero_centered_gamma) {
if (sizeof(InputType) < sizeof(OutputType)) {
GTEST_SKIP() << "LN kernel does not support OutputType > InputType";
return;
Expand DownExpand Up@@ -158,32 +168,34 @@ void performTest(const size_t N, const size_t H) {

// Forward kernel
float epsilon = 1e-5;
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto fwd_function = zero_centered_gamma ? nvte_layernorm1p_fwd : nvte_layernorm_fwd;
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
nvte_layernorm_fwd(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());
fwd_function(input.data(), gamma.data(), beta.data(), epsilon,
z.data(), mu.data(), rsigma.data(), 0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Backward kernel
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
auto bwd_function = zero_centered_gamma ? nvte_layernorm1p_bwd : nvte_layernorm_bwd;
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
workspace = Tensor(workspace.shape(), workspace.dtype());
barrier = Tensor(barrier.shape(), barrier.dtype());
dgamma_part = Tensor(dgamma_part.shape(), dgamma_part.dtype());
dbeta_part = Tensor(dbeta_part.shape(), dbeta_part.dtype());
nvte_layernorm_bwd(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());
bwd_function(dz.data(), input.data(),
mu.data(), rsigma.data(), gamma.data(),
dx.data(), dgamma.data(), dbeta.data(),
dgamma_part.data(), dbeta_part.data(),
0, prop.multiProcessorCount,
workspace.data(), barrier.data());

// Reference implementations
// use the GPU stats to tighten the tolerances
Expand All@@ -201,12 +213,13 @@ void performTest(const size_t N, const size_t H) {
rsigma.cpu_dptr<float>(),
N, H,
&ref_amax,
ref_scale);
ref_scale,
zero_centered_gamma);
compute_ref_backward(dz.cpu_dptr<WeightType>(), input.cpu_dptr<InputType>(),
mu.cpu_dptr<float>(), rsigma.cpu_dptr<float>(),
gamma.cpu_dptr<WeightType>(),
ref_dx.get(), ref_dgamma.get(), ref_dbeta.get(),
N, H);
N, H, zero_centered_gamma);

cudaDeviceSynchronize();
auto err = cudaGetLastError();
Expand DownExpand Up@@ -248,7 +261,8 @@ std::vector<std::pair<size_t, size_t>> test_cases = {{2048, 12288},

class LNTestSuite : public ::testing::TestWithParam<std::tuple<transformer_engine::DType,
transformer_engine::DType,
std::pair<size_t, size_t>>> {};
std::pair<size_t, size_t>,
bool>> {};

TEST_P(LNTestSuite, TestLN) {
using namespace transformer_engine;
Expand All@@ -257,10 +271,11 @@ TEST_P(LNTestSuite, TestLN) {
const DType input_type = std::get<0>(GetParam());
const DType output_type = std::get<1>(GetParam());
const auto size = std::get<2>(GetParam());
const bool zero_centered_gamma = std::get<3>(GetParam());

TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(input_type, InputType,
TRANSFORMER_ENGINE_TYPE_SWITCH_ALL(output_type, OutputType,
performTest<InputType, OutputType>(size.first, size.second);
performTest<InputType, OutputType>(size.first, size.second, zero_centered_gamma);
);
);
}
Expand All@@ -271,11 +286,13 @@ INSTANTIATE_TEST_SUITE_P(
::testing::Combine(
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16),
::testing::Values(DType::kFloat32, DType::kBFloat16, DType::kFloat16, DType::kFloat8E4M3),
::testing::ValuesIn(test_cases)),
::testing::ValuesIn(test_cases),
::testing::Values(false, true)),
[](const testing::TestParamInfo<LNTestSuite::ParamType>& info) {
std::string name = test::typeName(std::get<0>(info.param)) + "X" +
test::typeName(std::get<1>(info.param)) + "X" +
std::to_string(std::get<2>(info.param).first) + "X" +
std::to_string(std::get<2>(info.param).second);
std::to_string(std::get<2>(info.param).second) + "X" +
std::to_string(std::get<3>(info.param));
return name;
});
29 changes: 21 additions & 8 deletions tests/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -434,10 +434,12 @@ def forward(self, inp, weight):
@pytest.mark.parametrize("use_fp8", [False, True])
@pytest.mark.parametrize("scale_factor", [448, 112])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm(
use_fp8: bool,
scale_factor: float,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -459,7 +461,8 @@ def forward(self, inp):
inp,
self.weight,
self.bias,
self.eps)
self.eps,
zero_centered_gamma)
return ret

class TestFP8_Layernorm(nn.Module):
Expand All@@ -482,7 +485,8 @@ def forward(self, inp):
self.eps,
self.meta,
self.fp8_tensor,
self.fp8_type)
self.fp8_type,
zero_centered_gamma)

ret = cast_from_fp8(
ret,
Expand All@@ -500,7 +504,7 @@ def forward(self, inp):
do_export(model, inp, fname, use_fp8=use_fp8)
if precision not in (torch.bfloat16, ):
# TODO: FP32 has a small threshold (1e-5)
validate_result(fname, inp, model, atol=1e-3, is_fp8=use_fp8)
validate_result(fname, inp, model, atol=4e-3, is_fp8=use_fp8)


@skip_FP8
Expand DownExpand Up@@ -646,13 +650,15 @@ def forward(self, inp):
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_linear(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -676,6 +682,7 @@ def test_export_layernorm_linear(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand All@@ -698,13 +705,15 @@ def test_export_layernorm_linear(
(torch.float16, True),
(torch.float16, False),
])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_layernorm_mlp(
scale_factor: float,
use_fp8: bool,
use_bias: bool,
return_bias: bool,
return_layernorm_output: bool,
precision: torch.dtype
precision: torch.dtype,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand All@@ -729,6 +738,7 @@ def test_export_layernorm_mlp(
return_bias=return_bias,
return_layernorm_output=return_layernorm_output,
params_dtype=precision,
zero_centered_gamma=zero_centered_gamma,
).to(device='cuda')
if use_fp8:
set_layer_scale(model, scale_factor)
Expand DownExpand Up@@ -902,14 +912,16 @@ def test_export_multihead_attention(
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
use_mask: bool,
attn_mask_type: str,
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
if use_fp8 and torch.cuda.get_device_properties(torch.cuda.current_device()).major < 9:
Expand DownExpand Up@@ -947,7 +959,8 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling).to(device='cuda')
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
validate_result(fname, inp, model, atol=1e-3)
Expand Down
Loading