Skip to content
Prev Previous commit
Next Next commit
Add support for half for norm
  • Loading branch information
umar456 committed Jun 11, 2022
commit c388d51bcc3baf510b1c475117f975401b088c1c
65 changes: 33 additions & 32 deletions src/api/c/norm.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
#include <arith.hpp>
#include <backend.hpp>
#include <common/ArrayInfo.hpp>
#include <common/cast.hpp>
#include <common/err_common.hpp>
#include <complex.hpp>
#include <copy.hpp>
Expand All @@ -24,6 +25,7 @@
#include <af/traits.hpp>

using af::dim4;
using common::cast;
using detail::arithOp;
using detail::Array;
using detail::cdouble;
Expand All @@ -35,15 +37,21 @@ using detail::reduce;
using detail::reduce_all;
using detail::scalar;

template<typename T>
using normReductionResult =
typename std::conditional<std::is_same<T, common::half>::value, float,
T>::type;

template<typename T>
double matrixNorm(const Array<T> &A, double p) {
using RT = normReductionResult<T>;
if (p == 1) {
Array<T> colSum = reduce<af_add_t, T, T>(A, 0);
return getScalar<T>(reduce_all<af_max_t, T, T>(colSum));
Array<RT> colSum = reduce<af_add_t, T, normReductionResult<T>>(A, 0);
return getScalar<RT>(reduce_all<af_max_t, RT, RT>(colSum));
}
if (p == af::Inf) {
Array<T> rowSum = reduce<af_add_t, T, T>(A, 1);
return getScalar<T>(reduce_all<af_max_t, T, T>(rowSum));
Array<RT> rowSum = reduce<af_add_t, T, RT>(A, 1);
return getScalar<RT>(reduce_all<af_max_t, RT, RT>(rowSum));
}

AF_ERROR("This type of norm is not supported in ArrayFire\n",
Expand All @@ -52,41 +60,45 @@ double matrixNorm(const Array<T> &A, double p) {

template<typename T>
double vectorNorm(const Array<T> &A, double p) {
if (p == 1) { return getScalar<T>(reduce_all<af_add_t, T, T>(A)); }
using RT = normReductionResult<T>;
if (p == 1) { return getScalar<RT>(reduce_all<af_add_t, T, RT>(A)); }
if (p == af::Inf) {
return getScalar<T>(reduce_all<af_max_t, T, T>(A));
return getScalar<RT>(reduce_all<af_max_t, RT, RT>(cast<RT>(A)));
} else if (p == 2) {
Array<T> A_sq = arithOp<T, af_mul_t>(A, A, A.dims());
return std::sqrt(getScalar<T>(reduce_all<af_add_t, T, T>(A_sq)));
return std::sqrt(getScalar<RT>(reduce_all<af_add_t, T, RT>(A_sq)));
}

Array<T> P = createValueArray<T>(A.dims(), scalar<T>(p));
Array<T> A_p = arithOp<T, af_pow_t>(A, P, A.dims());
return std::pow(getScalar<T>(reduce_all<af_add_t, T, T>(A_p)), T(1.0 / p));
return std::pow(getScalar<RT>(reduce_all<af_add_t, T, RT>(A_p)), (1.0 / p));
}

template<typename T>
double LPQNorm(const Array<T> &A, double p, double q) {
Array<T> A_p_norm = createEmptyArray<T>(dim4());
using RT = normReductionResult<T>;
Array<RT> A_p_norm = createEmptyArray<RT>(dim4());

if (p == 1) {
A_p_norm = reduce<af_add_t, T, T>(A, 0);
A_p_norm = reduce<af_add_t, T, RT>(A, 0);
} else {
Array<T> P = createValueArray<T>(A.dims(), scalar<T>(p));
Array<T> invP = createValueArray<T>(A.dims(), scalar<T>(1.0 / p));
Array<T> P = createValueArray<T>(A.dims(), scalar<T>(p));
Array<RT> invP = createValueArray<RT>(A.dims(), scalar<RT>(1.0 / p));

Array<T> A_p = arithOp<T, af_pow_t>(A, P, A.dims());
Array<T> A_p_sum = reduce<af_add_t, T, T>(A_p, 0);
A_p_norm = arithOp<T, af_pow_t>(A_p_sum, invP, invP.dims());
Array<T> A_p = arithOp<T, af_pow_t>(A, P, A.dims());
Array<RT> A_p_sum = reduce<af_add_t, T, RT>(A_p, 0);
A_p_norm = arithOp<RT, af_pow_t>(A_p_sum, invP, invP.dims());
}

if (q == 1) { return getScalar<T>(reduce_all<af_add_t, T, T>(A_p_norm)); }
if (q == 1) {
return getScalar<RT>(reduce_all<af_add_t, RT, RT>(A_p_norm));
}

Array<T> Q = createValueArray<T>(A_p_norm.dims(), scalar<T>(q));
Array<T> A_p_norm_q = arithOp<T, af_pow_t>(A_p_norm, Q, Q.dims());
Array<RT> Q = createValueArray<RT>(A_p_norm.dims(), scalar<RT>(q));
Array<RT> A_p_norm_q = arithOp<RT, af_pow_t>(A_p_norm, Q, Q.dims());

return std::pow(getScalar<T>(reduce_all<af_add_t, T, T>(A_p_norm_q)),
T(1.0 / q));
return std::pow(getScalar<RT>(reduce_all<af_add_t, RT, RT>(A_p_norm_q)),
(1.0 / q));
}

template<typename T>
Expand All @@ -98,21 +110,13 @@ double norm(const af_array a, const af_norm_type type, const double p,

switch (type) {
case AF_NORM_EUCLID: return vectorNorm(A, 2);

case AF_NORM_VECTOR_1: return vectorNorm(A, 1);

case AF_NORM_VECTOR_INF: return vectorNorm(A, af::Inf);

case AF_NORM_VECTOR_P: return vectorNorm(A, p);

case AF_NORM_MATRIX_1: return matrixNorm(A, 1);

case AF_NORM_MATRIX_INF: return matrixNorm(A, af::Inf);

case AF_NORM_MATRIX_2: return matrixNorm(A, 2);

case AF_NORM_MATRIX_L_PQ: return LPQNorm(A, p, q);

default:
AF_ERROR("This type of norm is not supported in ArrayFire\n",
AF_ERR_NOT_SUPPORTED);
Expand All @@ -123,24 +127,21 @@ af_err af_norm(double *out, const af_array in, const af_norm_type type,
const double p, const double q) {
try {
const ArrayInfo &i_info = getInfo(in);

if (i_info.ndims() > 2) {
AF_ERROR("solve can not be used in batch mode", AF_ERR_BATCH);
}

af_dtype i_type = i_info.getType();

ARG_ASSERT(1, i_info.isFloating()); // Only floating and complex types

*out = 0;

if (i_info.ndims() == 0) { return AF_SUCCESS; }

switch (i_type) {
case f32: *out = norm<float>(in, type, p, q); break;
case f64: *out = norm<double>(in, type, p, q); break;
case c32: *out = norm<cfloat>(in, type, p, q); break;
case c64: *out = norm<cdouble>(in, type, p, q); break;
case f16: *out = norm<common::half>(in, type, p, q); break;
default: TYPE_ERROR(1, i_type);
}
}
Expand Down
72 changes: 40 additions & 32 deletions test/norm.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -19,18 +19,18 @@ using std::complex;
using std::stringstream;
using std::vector;

std::ostream& operator<<(std::ostream &os, af::normType nt) {
switch(nt) {
case AF_NORM_VECTOR_1: os << "AF_NORM_VECTOR_1"; break;
case AF_NORM_VECTOR_INF: os << "AF_NORM_VECTOR_INF"; break;
case AF_NORM_VECTOR_2: os << "AF_NORM_VECTOR_2"; break;
case AF_NORM_VECTOR_P: os << "AF_NORM_VECTOR_P"; break;
case AF_NORM_MATRIX_1: os << "AF_NORM_MATRIX_1"; break;
case AF_NORM_MATRIX_INF: os << "AF_NORM_MATRIX_INF"; break;
case AF_NORM_MATRIX_2: os << "AF_NORM_MATRIX_2"; break;
case AF_NORM_MATRIX_L_PQ: os << "AF_NORM_MATRIX_L_PQ"; break;
}
return os;
std::ostream &operator<<(std::ostream &os, af::normType nt) {
switch (nt) {
case AF_NORM_VECTOR_1: os << "AF_NORM_VECTOR_1"; break;
case AF_NORM_VECTOR_INF: os << "AF_NORM_VECTOR_INF"; break;
case AF_NORM_VECTOR_2: os << "AF_NORM_VECTOR_2"; break;
case AF_NORM_VECTOR_P: os << "AF_NORM_VECTOR_P"; break;
case AF_NORM_MATRIX_1: os << "AF_NORM_MATRIX_1"; break;
case AF_NORM_MATRIX_INF: os << "AF_NORM_MATRIX_INF"; break;
case AF_NORM_MATRIX_2: os << "AF_NORM_MATRIX_2"; break;
case AF_NORM_MATRIX_L_PQ: os << "AF_NORM_MATRIX_L_PQ"; break;
}
return os;
}

template<typename T>
Expand All @@ -51,12 +51,15 @@ double cpu_norm1_impl(af::dim4 &dims, std::vector<T> &value) {
double cpu_norm1(af::array &value) {
double norm1;
af::dim4 dims = value.dims();
if (value.type() == c32 || value.type() == c64) {
if (value.type() == f16) {
vector<half_float::half> values(value.elements());
value.host(values.data());
norm1 = cpu_norm1_impl<half_float::half>(dims, values);
} else if (value.type() == c32 || value.type() == c64) {
vector<complex<double> > values(value.elements());
value.as(c64).host(values.data());
norm1 = cpu_norm1_impl<complex<double> >(dims, values);
}
else {
} else {
vector<double> values(value.elements());
value.as(f64).host(values.data());
norm1 = cpu_norm1_impl<double>(dims, values);
Expand All @@ -71,7 +74,7 @@ double cpu_norm_inf_impl(af::dim4 &dims, std::vector<T> &value) {

double norm_inf = std::numeric_limits<double>::lowest();
for (int m = 0; m < M; m++) {
T *rowM = value.data() + m;
T *rowM = value.data() + m;
double sum = 0;
for (int n = 0; n < N; n++) { sum += abs(rowM[n * M]); }
norm_inf = std::max(norm_inf, sum);
Expand All @@ -95,16 +98,16 @@ double cpu_norm_inf(af::array &value) {
}

using norm_params = std::tuple<af::dim4, af::dtype>;
class Norm : public ::testing::TestWithParam<
std::tuple<af::dim4, af::dtype> > {};
class Norm
: public ::testing::TestWithParam<std::tuple<af::dim4, af::dtype> > {};

INSTANTIATE_TEST_CASE_P(
Norm, Norm,
::testing::Combine(::testing::Values(dim4(3, 3), dim4(32, 32), dim4(33, 33),
dim4(64, 64), dim4(128, 128),
dim4(129, 129), dim4(256, 256),
dim4(257, 257)),
::testing::Values(f32, f64, c32, c64)),
::testing::Values(f32, f64, c32, c64, f16)),
[](const ::testing::TestParamInfo<Norm::ParamType> info) {
stringstream ss;
using std::get;
Expand All @@ -116,48 +119,52 @@ INSTANTIATE_TEST_CASE_P(
TEST_P(Norm, Identity_AF_NORM_MATRIX_1) {
using std::get;
norm_params param = GetParam();
if (get<1>(param) == f16) SUPPORTED_TYPE_CHECK(half_float::half);
if (get<1>(param) == f64) SUPPORTED_TYPE_CHECK(double);

array identity = af::identity(get<0>(param), get<1>(param));
double result = norm(identity, AF_NORM_MATRIX_1);
double norm1 = cpu_norm1(identity);
double result = norm(identity, AF_NORM_MATRIX_1);
double norm1 = cpu_norm1(identity);

ASSERT_DOUBLE_EQ(norm1, result);
}

TEST_P(Norm, Random_AF_NORM_MATRIX_1) {
using std::get;
norm_params param = GetParam();
if (get<1>(param) == f16) SUPPORTED_TYPE_CHECK(half_float::half);
if (get<1>(param) == f64) SUPPORTED_TYPE_CHECK(double);

array identity = af::randu(get<0>(param), get<1>(param)) - 0.5f;
double result = norm(identity, AF_NORM_MATRIX_1);
double norm1 = cpu_norm1(identity);
array in = af::randu(get<0>(param), get<1>(param)) - 0.5f;
double result = norm(in, AF_NORM_MATRIX_1);
double norm1 = cpu_norm1(in);

ASSERT_NEAR(norm1, result, 2e-5);
ASSERT_NEAR(norm1, result, 2e-4);
}

TEST_P(Norm, Identity_AF_NORM_MATRIX_2_NOT_SUPPORTED) {
using std::get;
norm_params param = GetParam();
if (get<1>(param) == f16) SUPPORTED_TYPE_CHECK(half_float::half);
if (get<1>(param) == f64) SUPPORTED_TYPE_CHECK(double);
try {
double result =
norm(af::identity(get<0>(param), get<1>(param)), AF_NORM_MATRIX_2);
FAIL();
} catch (af::exception &ex) {
ASSERT_EQ(AF_ERR_NOT_SUPPORTED, ex.err());
return;
ASSERT_EQ(AF_ERR_NOT_SUPPORTED, ex.err());
return;
}
FAIL();
}

TEST_P(Norm, Identity_AF_NORM_MATRIX_INF) {
using std::get;
norm_params param = GetParam();
if (get<1>(param) == f16) SUPPORTED_TYPE_CHECK(half_float::half);
if (get<1>(param) == f64) SUPPORTED_TYPE_CHECK(double);
array in = af::identity(get<0>(param), get<1>(param));
double result = norm(in, AF_NORM_MATRIX_INF);
array in = af::identity(get<0>(param), get<1>(param));
double result = norm(in, AF_NORM_MATRIX_INF);
double norm_inf = cpu_norm_inf(in);

ASSERT_DOUBLE_EQ(norm_inf, result);
Expand All @@ -166,10 +173,11 @@ TEST_P(Norm, Identity_AF_NORM_MATRIX_INF) {
TEST_P(Norm, Random_AF_NORM_MATRIX_INF) {
using std::get;
norm_params param = GetParam();
if (get<1>(param) == f16) SUPPORTED_TYPE_CHECK(half_float::half);
if (get<1>(param) == f64) SUPPORTED_TYPE_CHECK(double);
array in = af::randu(get<0>(param), get<1>(param));
double result = norm(in, AF_NORM_MATRIX_INF);
array in = af::randu(get<0>(param), get<1>(param));
double result = norm(in, AF_NORM_MATRIX_INF);
double norm_inf = cpu_norm_inf(in);

ASSERT_NEAR(norm_inf, result, 2e-5);
ASSERT_NEAR(norm_inf, result, 2e-4);
}