Skip to content
This repository was archived by the owner on Aug 20, 2026. It is now read-only.
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
56 changes: 35 additions & 21 deletions reg_field.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -13,6 +13,20 @@ FF_NAMESPACE_BEGIN(FF)
FF_NAMESPACE_BEGIN(FF_DEVICE)
FF_NAMESPACE_BEGIN(reg_field)

//----------------------------------------------------------------------
// op dispatch helper (issue #6a)
//----------------------------------------------------------------------
// Thread the compile-time op ('=', '+', '-') through a function-template
// wrapper. Mirrors the CUDA impl's `Op<op,scalar_t,reduce_t>::f` dispatch, but
// as a named function template because C++11 (unlike the CUDA/C++17 path) only
// accepts a *function id-expression* -- not a constexpr pointer variable -- as
// a non-type template argument of function-pointer type.
template <char op, typename scalar_t, typename reduce_t>
inline scalar_t & op_apply(scalar_t & out, const reduce_t & in)
{
return Op<op, scalar_t, reduce_t>::f(out, in);
}

//======================================================================
// ABSOLUTE
//======================================================================
Expand DownExpand Up@@ -55,7 +69,7 @@ void matvec_absolute(
offset_t inp_offset = index2offset(i, nall, size, stride_inp);
offset_t out_offset = index2offset(i, nall, size, stride_out);

Impl::template matvec_absolute<set>(
Impl::template matvec_absolute<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, inp + inp_offset, osc, isc, kernel, nc);
}});
delete[] kernel;
Expand DownExpand Up@@ -98,7 +112,7 @@ void kernel_absolute(
offset_t out_offset = index2offset(i, nbatch, size, stride);
out_offset += offset;

Impl::template kernel_absolute<set>(out + out_offset, sc, kernel, nc);
Impl::template kernel_absolute<op_apply<op, scalar_t, reduce_t> >(out + out_offset, sc, kernel, nc);
}});
delete[] kernel;
}
Expand DownExpand Up@@ -138,7 +152,7 @@ void diag_absolute(
offset_t loc[ndim];
offset_t out_offset = index2offset_v2<ndim>(i, nall, size, stride, loc);

Impl::template diag_absolute<set>(out + out_offset, sc, kernel, nc);
Impl::template diag_absolute<op_apply<op, scalar_t, reduce_t> >(out + out_offset, sc, kernel, nc);
}});
delete[] kernel;
}
Expand DownExpand Up@@ -189,7 +203,7 @@ void matvec_membrane(
offset_t inp_offset = index2offset_v2<ndim>(i, nall, size, stride_inp, loc);
offset_t out_offset = index2offset(i, nall, size, stride_out);

Impl::template matvec_membrane<set>(
Impl::template matvec_membrane<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, inp + inp_offset,
loc, size + nbatch, stride_inp + nbatch, osc, isc, kernel, nc);
}});
Expand DownExpand Up@@ -236,7 +250,7 @@ void kernel_membrane(
offset_t out_offset = index2offset(i, nbatch, size, stride);
out_offset += offset;

Impl::template kernel_membrane<set>(
Impl::template kernel_membrane<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, sc, stride + nbatch, kernel, nc);
}});
delete[] kernel;
Expand DownExpand Up@@ -280,7 +294,7 @@ void diag_membrane(
offset_t loc[ndim];
offset_t out_offset = index2offset_v2<ndim>(i, nall, size, stride, loc);

Impl::template diag_membrane<set>(
Impl::template diag_membrane<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, sc, loc, size + nbatch, kernel, nc);
}});
delete[] kernel;
Expand DownExpand Up@@ -419,7 +433,7 @@ void matvec_bending(
offset_t inp_offset = index2offset_v2<ndim>(i, nall, size, stride_inp, loc);
offset_t out_offset = index2offset(i, nall, size, stride_out);

Impl::template matvec_bending<set>(
Impl::template matvec_bending<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, inp + inp_offset,
loc, size + nbatch, stride_inp + nbatch, osc, isc, kernel, nc);
}
Expand DownExpand Up@@ -468,7 +482,7 @@ void kernel_bending(
offset_t out_offset = index2offset(i, nbatch, size, stride);
out_offset += offset;

Impl::template kernel_bending<set>(
Impl::template kernel_bending<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, sc, stride + nbatch, kernel, nc);
}});
delete[] kernel;
Expand DownExpand Up@@ -513,7 +527,7 @@ void diag_bending(
offset_t loc[ndim];
offset_t out_offset = index2offset_v2<ndim>(i, nall, size, stride, loc);

Impl::template diag_bending<set>(
Impl::template diag_bending<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, sc, loc, size + nbatch, kernel, nc);
}});
delete[] kernel;
Expand DownExpand Up@@ -650,7 +664,7 @@ void matvec_absolute_rls(
offset_t out_offset = index2offset(i, nall, size, stride_out);
offset_t wgt_offset = index2offset(i, nall, size, stride_wgt);

Impl::template matvec_absolute_rls<set>(
Impl::template matvec_absolute_rls<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, inp + inp_offset, wgt + wgt_offset,
osc, isc, wsc, kernel, nc);
}});
Expand DownExpand Up@@ -694,7 +708,7 @@ void diag_absolute_rls(
offset_t out_offset = index2offset(i, nall, size, stride_out);
offset_t wgt_offset = index2offset(i, nall, size, stride_wgt);

Impl::template diag_absolute_rls<set>(
Impl::template diag_absolute_rls<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, wgt + wgt_offset, osc, wsc, kernel, nc);
}});
delete[] kernel;
Expand DownExpand Up@@ -827,7 +841,7 @@ void matvec_absolute_jrls(
offset_t out_offset = index2offset(i, nall, size, stride_out);
offset_t wgt_offset = index2offset(i, nall, size, stride_wgt);

Impl::template matvec_absolute_jrls<set>(
Impl::template matvec_absolute_jrls<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, inp + inp_offset, wgt + wgt_offset,
osc, isc, kernel, nc);
}});
Expand DownExpand Up@@ -870,7 +884,7 @@ void diag_absolute_jrls(
offset_t out_offset = index2offset(i, nall, size, stride_out);
offset_t wgt_offset = index2offset(i, nall, size, stride_wgt);

Impl::template diag_absolute_jrls<set>(
Impl::template diag_absolute_jrls<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, wgt + wgt_offset, osc, kernel, nc);
}});
delete[] kernel;
Expand DownExpand Up@@ -1008,7 +1022,7 @@ void matvec_membrane_rls(
offset_t out_offset = index2offset(i, nall, size, stride_out);
offset_t wgt_offset = index2offset(i, nall, size, stride_wgt);

Impl::template matvec_membrane_rls<set>(
Impl::template matvec_membrane_rls<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, inp + inp_offset, wgt + wgt_offset,
loc, size + nbatch, stride_inp + nbatch, stride_wgt + nbatch,
osc, isc, wsc, kernel, nc);
Expand DownExpand Up@@ -1058,7 +1072,7 @@ void diag_membrane_rls(
offset_t out_offset = index2offset_v2<ndim>(i, nall, size, stride_out, loc);
offset_t wgt_offset = index2offset(i, nall, size, stride_wgt);

Impl::template diag_membrane_rls<set>(
Impl::template diag_membrane_rls<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, wgt + wgt_offset,
loc, size + nbatch, stride_wgt + nbatch, osc, wsc, kernel, nc);
}});
Expand DownExpand Up@@ -1207,7 +1221,7 @@ void matvec_membrane_jrls(
offset_t out_offset = index2offset(i, nall, size, stride_out);
offset_t wgt_offset = index2offset(i, nall, size, stride_wgt);

Impl::template matvec_membrane_jrls<set>(
Impl::template matvec_membrane_jrls<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, inp + inp_offset, wgt + wgt_offset,
loc, size + nbatch, stride_inp + nbatch, stride_wgt + nbatch,
osc, isc, kernel, nc);
Expand DownExpand Up@@ -1256,7 +1270,7 @@ void diag_membrane_jrls(
offset_t out_offset = index2offset_v2<ndim>(i, nall, size, stride_out, loc);
offset_t wgt_offset = index2offset(i, nall, size, stride_wgt);

Impl::template diag_membrane_jrls<set>(
Impl::template diag_membrane_jrls<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, wgt + wgt_offset,
loc, size + nbatch, stride_wgt + nbatch, osc, kernel, nc);
}});
Expand DownExpand Up@@ -1405,7 +1419,7 @@ void matvec_bending_rls(
offset_t out_offset = index2offset(i, nall, size, stride_out);
offset_t wgt_offset = index2offset(i, nall, size, stride_wgt);

Impl::template matvec_bending_rls<set>(
Impl::template matvec_bending_rls<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, inp + inp_offset, wgt + wgt_offset,
loc, size + nbatch, stride_inp + nbatch, stride_wgt + nbatch,
osc, isc, wsc, kernel, nc);
Expand DownExpand Up@@ -1456,7 +1470,7 @@ void diag_bending_rls(
offset_t out_offset = index2offset_v2<ndim>(i, nall, size, stride_out, loc);
offset_t wgt_offset = index2offset(i, nall, size, stride_wgt);

Impl::template diag_bending_rls<set>(
Impl::template diag_bending_rls<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, wgt + wgt_offset,
loc, size + nbatch, stride_wgt + nbatch, osc, wsc, kernel, nc);
}});
Expand DownExpand Up@@ -1606,7 +1620,7 @@ void matvec_bending_jrls(
offset_t out_offset = index2offset(i, nall, size, stride_out);
offset_t wgt_offset = index2offset(i, nall, size, stride_wgt);

Impl::template matvec_bending_jrls<set>(
Impl::template matvec_bending_jrls<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, inp + inp_offset, wgt + wgt_offset,
loc, size + nbatch, stride_inp + nbatch, stride_wgt + nbatch,
osc, isc, kernel, nc);
Expand DownExpand Up@@ -1656,7 +1670,7 @@ void diag_bending_jrls(
offset_t out_offset = index2offset_v2<ndim>(i, nall, size, stride_out, loc);
offset_t wgt_offset = index2offset(i, nall, size, stride_wgt);

Impl::template diag_bending_jrls<set>(
Impl::template diag_bending_jrls<op_apply<op, scalar_t, reduce_t> >(
out + out_offset, wgt + wgt_offset,
loc, size + nbatch, stride_wgt + nbatch, osc, kernel, nc);
}});
Expand Down
Loading