Update hyperparameter_validation.cpp

This commit is contained in:
bjyb 2023-09-18 20:20:37 +08:00
parent 4387e8852a
commit 1fbc036907
1 changed files with 137 additions and 3 deletions

View File

@ -28,6 +28,8 @@
#include "nodes/plannodes.h" #include "nodes/plannodes.h"
#include "db4ai/db4ai_api.h" #include "db4ai/db4ai_api.h"
//Used to add, delete, check and modify the value of the superparameter.
/////////////////////////////////////////////////////////////////////////////// ///////////////////////////////////////////////////////////////////////////////
@ -37,6 +39,10 @@
is_supervised \ is_supervised \
} }
/*Function: get_ Hyperparameter_ Definitions
Parameter: (AlgorithmML algorithm, int32_t * result_size)
Return: HyperparameterDefinition*
Enter the algorithm and number of result digits to return the definition of hyperparameters for this model.*/
const HyperparameterDefinition* get_hyperparameter_definitions(AlgorithmML algorithm, int32_t *result_size) const HyperparameterDefinition* get_hyperparameter_definitions(AlgorithmML algorithm, int32_t *result_size)
{ {
@ -44,6 +50,11 @@ const HyperparameterDefinition* get_hyperparameter_definitions(AlgorithmML algor
return api->get_hyperparameters_definitions(api, result_size); return api->get_hyperparameters_definitions(api, result_size);
} }
/*Function: get_ Algorithm_ Configuration
Parameter: AlgorithmML algorithm
Return: AlgorithmConfiguration*
Determine if the algorithm exists.*/
AlgorithmConfiguration *get_algorithm_configuration(AlgorithmML algorithm) AlgorithmConfiguration *get_algorithm_configuration(AlgorithmML algorithm)
{ {
switch (algorithm) { switch (algorithm) {
@ -55,6 +66,16 @@ AlgorithmConfiguration *get_algorithm_configuration(AlgorithmML algorithm)
return NULL; return NULL;
} }
/*
Function: find_hyperparameter_definition
Parameter:(const HyperparameterDefinition definitions[],
int32_t definitions_size,
const char *hyperparameter_name)
Return:HyperparameterDefinition *
Enter the model name and return the model definition.
*/
const HyperparameterDefinition *find_hyperparameter_definition(const HyperparameterDefinition definitions[], const HyperparameterDefinition *find_hyperparameter_definition(const HyperparameterDefinition definitions[],
int32_t definitions_size, int32_t definitions_size,
const char *hyperparameter_name) const char *hyperparameter_name)
@ -68,6 +89,14 @@ const HyperparameterDefinition *find_hyperparameter_definition(const Hyperparame
} }
// Set the value of a hyperparameter structure // Set the value of a hyperparameter structure
/*
Function: set_ Hyperparameter_ Datum
Parameter: (Hyperparameter * hyperp, Oid type, Datum value)
Return: None
Set the hyperparameter type and value, which is called by the system.
An exception is thrown if there is no super parameter.
*/
static void set_hyperparameter_datum(Hyperparameter *hyperp, Oid type, Datum value) static void set_hyperparameter_datum(Hyperparameter *hyperp, Oid type, Datum value)
{ {
if (type == ANYENUMOID) { // Outside of hyperparameter module, treat them as strings if (type == ANYENUMOID) { // Outside of hyperparameter module, treat them as strings
@ -80,6 +109,14 @@ static void set_hyperparameter_datum(Hyperparameter *hyperp, Oid type, Datum val
} }
} }
/*
Function: add_model_hyperparameter
Parameter: (List *hyperparameters, MemoryContext memcxt, const char *name, Oid type,
Datum value)
Return: List *
Adding hyperparameters to the model
*/
static List *add_model_hyperparameter(List *hyperparameters, MemoryContext memcxt, const char *name, Oid type, static List *add_model_hyperparameter(List *hyperparameters, MemoryContext memcxt, const char *name, Oid type,
Datum value) Datum value)
{ {
@ -93,6 +130,13 @@ static List *add_model_hyperparameter(List *hyperparameters, MemoryContext memcx
return hyperparameters; return hyperparameters;
} }
/*
Function: update_model_hyperparameter
Parameter: (MemoryContext memcxt, List *hyperparameters, const char *name, Oid type, Datum value)
Return:None
update hyperparameters to the model
*/
void update_model_hyperparameter(MemoryContext memcxt, List *hyperparameters, const char *name, Oid type, Datum value) void update_model_hyperparameter(MemoryContext memcxt, List *hyperparameters, const char *name, Oid type, Datum value)
{ {
@ -108,12 +152,20 @@ void update_model_hyperparameter(MemoryContext memcxt, List *hyperparameters, co
MemoryContextSwitchTo(old_context); MemoryContextSwitchTo(old_context);
} }
/* inline change bool to str*/
inline const char *bool_to_str(bool value) inline const char *bool_to_str(bool value)
{ {
return value ? "TRUE" : "FALSE"; return value ? "TRUE" : "FALSE";
} }
/*Function: ereport_ Hyperparameter
Formal parameters: (int level, const char * name, Datum value, Oid type)
Return value: None
Display model hyperparameters*/
static void ereport_hyperparameter(int level, const char *name, Datum value, Oid type) static void ereport_hyperparameter(int level, const char *name, Datum value, Oid type)
{ {
switch (type) { switch (type) {
@ -214,6 +266,15 @@ static Datum get_hyperparameter(const HyperparameterDefinition *definition, void
// Set hyperparameter in hyperparameter struct to the givne value in the datum. Definition is used for metadata // Set hyperparameter in hyperparameter struct to the givne value in the datum. Definition is used for metadata
/*Function: set_ Hyperparameter
Formal parameters: (const HyperparameterDefinition * definition, Datum value, void * hyperparameter_struct)
Return value: None
Modify model hyperparameter values*/
static void set_hyperparameter(const HyperparameterDefinition *definition, Datum value, void *hyperparameter_struct) static void set_hyperparameter(const HyperparameterDefinition *definition, Datum value, void *hyperparameter_struct)
{ {
switch (definition->type) { switch (definition->type) {
@ -259,6 +320,17 @@ static void set_hyperparameter(const HyperparameterDefinition *definition, Datum
} }
} }
/*
Function: validate_ Hyperparameter_ String
Formal parameters: (const char * name, const char * value, const char * valid_values [],
Int32_ T valid_ Values_ Size)
Return value: None
Given the hyperparameter name, modify the model hyperparameter value.*/
static void validate_hyperparameter_string(const char *name, const char *value, const char *valid_values[], static void validate_hyperparameter_string(const char *name, const char *value, const char *valid_values[],
int32_t valid_values_size) int32_t valid_values_size)
{ {
@ -284,6 +356,11 @@ static void validate_hyperparameter_string(const char *name, const char *value,
} }
} }
/*Function: validate_ Hyperparameter
Formal parameters: (Datum value, Oid type, const HyperparameterValidation * validation, const char * name)
Return value: None
Make the modified hyperparameter values effective.*/
static void validate_hyperparameter(Datum value, Oid type, const HyperparameterValidation *validation, const char *name) static void validate_hyperparameter(Datum value, Oid type, const HyperparameterValidation *validation, const char *name)
{ {
switch (type) { switch (type) {
@ -354,6 +431,16 @@ static void validate_hyperparameter(Datum value, Oid type, const HyperparameterV
} }
} }
/*Function: extract_ Value_ From_ Variable_ Set_ Stmt
Parameter: (VariableSetStmt * stmt)
Return value: Value
Obtain modified values using preprocessing.*/
/*STMT is a C API provided by MySQL,
which is used to execute Prepared statements.
Compared with the direct execution of SQL, the
preprocessing statement has higher running
efficiency and better security.*/
static Value *extract_value_from_variable_set_stmt(VariableSetStmt *stmt) static Value *extract_value_from_variable_set_stmt(VariableSetStmt *stmt)
{ {
if (list_length(stmt->args) > 1) { if (list_length(stmt->args) > 1) {
@ -370,6 +457,14 @@ static Value *extract_value_from_variable_set_stmt(VariableSetStmt *stmt)
return value; return value;
} }
/*Function: value_ To_ Datum
Formal parameters: (Value * value, Oid expected_type, const char * name)
Return value: Datum
Modify the hyperparameter value to Datum type.*/
static Datum value_to_datum(Value *value, Oid expected_type, const char *name) static Datum value_to_datum(Value *value, Oid expected_type, const char *name)
{ {
Datum result = (Datum)0; Datum result = (Datum)0;
@ -457,6 +552,10 @@ static Datum value_to_datum(Value *value, Oid expected_type, const char *name)
return result; return result;
} }
/*Function: extract_ Datum_ From_ Variable_ Set_ Stmt
Formal parameters: (VariableSetStmt * stmt, const HyperparameterDefinition * definition)
Return value: Datum
Use preprocessing to obtain and modify Datum.*/
Datum extract_datum_from_variable_set_stmt(VariableSetStmt *stmt, const HyperparameterDefinition *definition) Datum extract_datum_from_variable_set_stmt(VariableSetStmt *stmt, const HyperparameterDefinition *definition)
{ {
Datum selected_value = (Datum)0; Datum selected_value = (Datum)0;
@ -470,6 +569,16 @@ Datum extract_datum_from_variable_set_stmt(VariableSetStmt *stmt, const Hyperpar
return selected_value; return selected_value;
} }
/*Function: configure_ Hyperparameters_ VSET
Formal parameters: (const HyperparameterDefinition definitions [], int32_t definitions_size,
List * hyperparameters, void * configuration)
Return value: Datum
Initialize hyperparameter configuration using set.*/
void configure_hyperparameters_vset(const HyperparameterDefinition definitions[], int32_t definitions_size, void configure_hyperparameters_vset(const HyperparameterDefinition definitions[], int32_t definitions_size,
List *hyperparameters, void *configuration) List *hyperparameters, void *configuration)
{ {
@ -543,6 +652,16 @@ void configure_hyperparameters(const HyperparameterDefinition definitions[], int
} }
} }
/*Function: prepare_ Model_ Hyperparameters
Formal parameters: (const HyperparameterDefinition * definitions, int32_t definitions_size,
Void * hyperparameter_ Struct, MemoryContext memcxt)
Return value: List*
Prepare model hyperparameters.*/
List *prepare_model_hyperparameters(const HyperparameterDefinition *definitions, int32_t definitions_size, List *prepare_model_hyperparameters(const HyperparameterDefinition *definitions, int32_t definitions_size,
void *hyperparameter_struct, MemoryContext memcxt) void *hyperparameter_struct, MemoryContext memcxt)
{ {
@ -555,6 +674,17 @@ List *prepare_model_hyperparameters(const HyperparameterDefinition *definitions,
return hyperparameters; return hyperparameters;
} }
/*Function: init_ Hyperparameters_ With_ Defaults
Formal parameters: (const HyperparameterDefinition definitions [], int32_t definitions_size,
Void * hyperparameter_ Struct
Return value: None
Initialize hyperparameters
*/
void init_hyperparameters_with_defaults(const HyperparameterDefinition definitions[], int32_t definitions_size, void init_hyperparameters_with_defaults(const HyperparameterDefinition definitions[], int32_t definitions_size,
void *hyperparameter_struct) void *hyperparameter_struct)
{ {
@ -563,6 +693,10 @@ void init_hyperparameters_with_defaults(const HyperparameterDefinition definitio
} }
} }
/*Function: print_ Hyperparameters
Formal parameters: (int level, List * hyperparameters)
Return value: None
Output all hyperparameter attributes*/
void print_hyperparameters(int level, List *hyperparameters) void print_hyperparameters(int level, List *hyperparameters)
{ {
foreach_cell(it, hyperparameters) { foreach_cell(it, hyperparameters) {