Skip to content
Open
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
38 changes: 13 additions & 25 deletions src/shamrock/include/shamrock/patch/FieldVariant.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@

#include "shambase/exception.hpp"
#include "shambackends/typeAliasVec.hpp"
#include <utility>
#include <variant>

namespace shamrock::patch {
Expand Down Expand Up @@ -65,49 +66,36 @@ namespace shamrock::patch {
"the type asked is not correct");
}

/**
* @brief Apply a visitor to the held alternative.
*
* The functor is forwarded to std::visit unchanged. Wrapping it in an extra
* generic lambda would give every call site a unique visitor type, and
* libstdc++ builds a visit vtable for each of those types.
*/
template<class Func>
void visit(Func &&f) {
std::visit(
[&](auto &arg) {
f(arg);
},
value);
std::visit(std::forward<Func>(f), value);
}

template<class Func>
auto visit_return(Func &&f) {
return std::visit(
[&](auto &arg) {
return f(arg);
},
value);
return std::visit(std::forward<Func>(f), value);
}

template<class Func>
void visit(Func &&f) const {
std::visit(
[&](auto &arg) {
f(arg);
},
value);
std::visit(std::forward<Func>(f), value);
}

template<class Func>
auto visit_return(Func &&f) const {
return std::visit(
[&](auto &arg) {
return f(arg);
},
value);
return std::visit(std::forward<Func>(f), value);
}

template<template<class> class Container2, class Func>
FieldVariant<Container2> convert(Func &&f) {
return std::visit(
[&](auto &arg) {
return f(arg);
},
value);
return std::visit(std::forward<Func>(f), value);
}
};

Expand Down
114 changes: 10 additions & 104 deletions src/shamrock/include/shamrock/patch/PatchDataLayer.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -79,19 +79,7 @@ namespace shamrock::patch {
init_fields();
}

inline PatchDataLayer(const PatchDataLayer &other) : pdl_ptr(other.get_layout_ptr()) {

NamedStackEntry stack_loc{"PatchDataLayer::copy_constructor", true};

for (auto &field_var : other.fields) {

field_var.visit([&](auto &field) {
using base_t =
typename std::remove_reference<decltype(field)>::type::Field_type;
fields.emplace_back(PatchDataField<base_t>(field));
});
};
}
PatchDataLayer(const PatchDataLayer &other);

/**
* @brief PatchDataLayer move constructor
Expand Down Expand Up @@ -120,18 +108,14 @@ namespace shamrock::patch {
template<class Functor>
inline void for_each_field_any(Functor &&func) {
for (auto &f : fields) {
f.visit([&](auto &arg) {
func(arg);
});
f.visit(func);
}
}

template<class Functor>
inline void for_each_field_any(Functor &&func) const {
for (auto &f : fields) {
f.visit([&](const auto &arg) {
func(arg);
});
f.visit(func);
}
}

Expand Down Expand Up @@ -234,42 +218,13 @@ namespace shamrock::patch {
void append_subset_to(
const sham::DeviceBuffer<u32> &idxs_buf, u32 sz, PatchDataLayer &pdat) const;

inline u32 get_obj_cnt() const {
u32 get_obj_cnt() const;

bool is_empty = fields.empty();

if (!is_empty) {
return fields[0].visit_return([](const auto &field) {
return field.get_obj_cnt();
});
}

throw shambase::make_except_with_loc<std::runtime_error>(
"this PatchDataLayer does not contain any fields");
}

inline u64 memsize() {
u64 sum = 0;

for (auto &field_var : fields) {

field_var.visit([&](auto &field) {
sum += field.memsize();
});
}

return sum;
}
u64 memsize();

inline bool is_empty() { return get_obj_cnt() == 0; }

void synchronize_buf() {
for (auto &field_var : fields) {
field_var.visit([&](auto &field) {
field.synchronize_buf();
});
}
}
void synchronize_buf();

void overwrite(PatchDataLayer &pdat, u32 obj_cnt);

Expand Down Expand Up @@ -389,17 +344,7 @@ namespace shamrock::patch {
* @brief check that all contained field have the same obj cnt
*
*/
inline void check_field_obj_cnt_match() {
u32 cnt = get_obj_cnt();
for (auto &field_var : fields) {
field_var.visit([&](auto &field) {
if (field.get_obj_cnt() != cnt) {
throw shambase::make_except_with_loc<std::runtime_error>(
"mismatch in obj cnt");
}
});
}
}
void check_field_obj_cnt_match();

// template<class T> inline std::vector<PatchDataField<T> & > get_field_list(){
// std::vector<PatchDataField<T> & > ret;
Expand Down Expand Up @@ -431,48 +376,9 @@ namespace shamrock::patch {

void fields_raz();

bool has_nan() {
StackEntry stack_loc{};

bool ret = false;

for (auto &field_var : fields) {
field_var.visit([&](auto &field) {
if (field.has_nan()) {
ret = true;
}
});
}
return ret;
}
bool has_inf() {
StackEntry stack_loc{};

bool ret = false;

for (auto &field_var : fields) {
field_var.visit([&](auto &field) {
if (field.has_inf()) {
ret = true;
}
});
}
return ret;
}
bool has_nan_or_inf() {
StackEntry stack_loc{};

bool ret = false;

for (auto &field_var : fields) {
field_var.visit([&](auto &field) {
if (field.has_nan_or_inf()) {
ret = true;
}
});
}
return ret;
}
bool has_nan();
bool has_inf();
bool has_nan_or_inf();

/**
* @brief
Expand Down
32 changes: 11 additions & 21 deletions src/shamrock/include/shamrock/patch/PatchDataLayerLayout.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
#include "shambase/exception.hpp"
#include "shambase/string.hpp"
#include "nlohmann/json_fwd.hpp"
#include "shamrock/legacy/patch/base/enabled_fields.hpp"
#include "shamrock/patch/FieldVariant.hpp"
#include "shamsys/legacy/log.hpp"
#include <sstream>
Expand Down Expand Up @@ -187,9 +188,7 @@ namespace shamrock::patch {
template<class Functor>
inline void for_each_field_any(Functor &&func) const {
for (auto &f : fields) {
f.visit([&](auto &arg) {
func(arg);
});
f.visit(func);
}
}

Expand Down Expand Up @@ -296,24 +295,15 @@ namespace shamrock::patch {
// out of line implementation of the PatchDataLayerLayout
////////////////////////////////////////////////////////////////////////////////////////////////

template<class T>
inline void PatchDataLayerLayout::add_field(
const std::string &field_name, u32 nvar, SourceLocation loc) {
if (has_field_name(field_name)) {
throw shambase::make_except_with_loc<std::invalid_argument>(
"add_field -> the name already exists");
}

shamlog_debug_ln(
"PatchDataLayerLayout",
"adding field :",
field_name,
nvar,
"loc :",
loc.format_one_line());

fields.push_back(var_t{FieldDescriptor<T>(field_name, nvar)});
}
#ifndef DOXYGEN
// Explicit instantiations live in PatchDataLayerLayout.cpp so TUs that call add_field
// do not instantiate FieldVariant construction for every enabled field type.
#define X(type) \
extern template void PatchDataLayerLayout::add_field<type>( \
const std::string &field_name, u32 nvar, SourceLocation loc);
XMAC_LIST_ENABLED_FIELD
#undef X
#endif

template<class T>
inline PatchDataLayerLayout::FieldDescriptor<T> PatchDataLayerLayout::get_field(
Expand Down
102 changes: 102 additions & 0 deletions src/shamrock/src/patch/PatchDataLayer.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,108 @@ namespace shamrock::patch {
});
}

PatchDataLayer::PatchDataLayer(const PatchDataLayer &other) : pdl_ptr(other.get_layout_ptr()) {

NamedStackEntry stack_loc{"PatchDataLayer::copy_constructor", true};

for (auto &field_var : other.fields) {

field_var.visit([&](auto &field) {
using base_t = typename std::remove_reference<decltype(field)>::type::Field_type;
fields.emplace_back(PatchDataField<base_t>(field));
});
}
}

u32 PatchDataLayer::get_obj_cnt() const {

if (!fields.empty()) {
return fields[0].visit_return([](const auto &field) {
return field.get_obj_cnt();
});
}

throw shambase::make_except_with_loc<std::runtime_error>(
"this PatchDataLayer does not contain any fields");
}

u64 PatchDataLayer::memsize() {
u64 sum = 0;

for (auto &field_var : fields) {

field_var.visit([&](auto &field) {
sum += field.memsize();
});
}

return sum;
}

void PatchDataLayer::synchronize_buf() {
for (auto &field_var : fields) {
field_var.visit([&](auto &field) {
field.synchronize_buf();
});
}
}

void PatchDataLayer::check_field_obj_cnt_match() {
u32 cnt = get_obj_cnt();
for (auto &field_var : fields) {
field_var.visit([&](auto &field) {
if (field.get_obj_cnt() != cnt) {
throw shambase::make_except_with_loc<std::runtime_error>("mismatch in obj cnt");
}
});
}
}

bool PatchDataLayer::has_nan() {
StackEntry stack_loc{};

bool ret = false;

for (auto &field_var : fields) {
field_var.visit([&](auto &field) {
if (field.has_nan()) {
ret = true;
}
});
}
return ret;
}

bool PatchDataLayer::has_inf() {
StackEntry stack_loc{};

bool ret = false;

for (auto &field_var : fields) {
field_var.visit([&](auto &field) {
if (field.has_inf()) {
ret = true;
}
});
}
return ret;
}

bool PatchDataLayer::has_nan_or_inf() {
StackEntry stack_loc{};

bool ret = false;

for (auto &field_var : fields) {
field_var.visit([&](auto &field) {
if (field.has_nan_or_inf()) {
ret = true;
}
});
}
return ret;
}

void PatchDataLayer::extract_element(u32 pidx, PatchDataLayer &out_pdat) {
StackEntry stack_loc{};

Expand Down
Loading
Loading