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 "db4ai/db4ai_api.h"
//Used to add, delete, check and modify the value of the superparameter.
///////////////////////////////////////////////////////////////////////////////
@ -37,13 +39,22 @@
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)
{
AlgorithmAPI* api = get_algorithm_api(algorithm);
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)
{
switch (algorithm) {
@ -55,6 +66,16 @@ AlgorithmConfiguration *get_algorithm_configuration(AlgorithmML algorithm)
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[],
int32_t definitions_size,
const char *hyperparameter_name)
@ -68,6 +89,14 @@ const HyperparameterDefinition *find_hyperparameter_definition(const Hyperparame
}
// 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)
{
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,
Datum value)
{
@ -93,7 +130,14 @@ static List *add_model_hyperparameter(List *hyperparameters, MemoryContext memcx
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)
{
MemoryContext old_context = MemoryContextSwitchTo(memcxt);
@ -108,12 +152,20 @@ void update_model_hyperparameter(MemoryContext memcxt, List *hyperparameters, co
MemoryContextSwitchTo(old_context);
}
/* inline change bool to str*/
inline const char *bool_to_str(bool value)
{
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)
{
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
/*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)
{
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[],
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)
{
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)
{
if (list_length(stmt->args) > 1) {
@ -370,6 +457,14 @@ static Value *extract_value_from_variable_set_stmt(VariableSetStmt *stmt)
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)
{
Datum result = (Datum)0;
@ -457,6 +552,10 @@ static Datum value_to_datum(Value *value, Oid expected_type, const char *name)
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 selected_value = (Datum)0;
@ -470,6 +569,16 @@ Datum extract_datum_from_variable_set_stmt(VariableSetStmt *stmt, const Hyperpar
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,
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,
void *hyperparameter_struct, MemoryContext memcxt)
{
@ -555,6 +674,17 @@ List *prepare_model_hyperparameters(const HyperparameterDefinition *definitions,
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 *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)
{
foreach_cell(it, hyperparameters) {