#ifndef GEMMSTONE_INCLUDE_GEMMSTONE_MICROKERNEL_PROTOCOL_HPP
#define GEMMSTONE_INCLUDE_GEMMSTONE_MICROKERNEL_PROTOCOL_HPP
#include "gemmstone/config.hpp"
GEMMSTONE_NAMESPACE_START
namespace microkernel {
struct StructuredType {
enum Type { u64,
s64,
u32,
s32,
u16,
s16,
u8,
s8,
u4,
s4, f64,
f32,
f16,
bf16,
bf8,
hf8,
f8_e8m0,
f4_e2m1,
f4_e3m0, any, } type
= Type::any;
enum Format { Scalar, GlobalPointer, LocalPointer, Tensor } format = Scalar;
int ndims = 1;
StructuredType() = default;
StructuredType(Type type_) : type(type_) {}
StructuredType(Format format_) : format(format_) {}
StructuredType(int ndims_) : format(Tensor), ndims(ndims_) {}
};
class Protocol {
public:
struct Argument {
const char *name;
enum { In = 0b01, Out = 0b10, InOut = In | Out } direction;
StructuredType stype;
bool in() const { return direction & In; }
bool out() const { return direction & Out; }
};
struct Setting {
const char *name;
};
Protocol() = default;
Protocol(std::string name, std::vector<Argument> arguments,
std::vector<Setting> settings)
: kernelBaseName_(std::move(name)), arguments_(std::move(arguments)), settings_(std::move(settings)) {}
const std::string &kernelBaseName() const { return kernelBaseName_; }
const std::vector<Argument> &arguments() const {
return arguments_;
}
const std::vector<Setting> &settings() const { return settings_; }
private:
std::string kernelBaseName_;
std::vector<Argument> arguments_;
std::vector<Setting> settings_;
};
}
GEMMSTONE_NAMESPACE_END
#endif