mirror of
https://github.com/NicolasBohn/NexQuant.git
synced 2026-07-28 07:57:44 +00:00
dbe2cf12bb
* Init todo * Evaluation & dataset * Generate new data * dataset generation * add the result * Analysis * Factor update * Updates * Reformat analysis.py * CI fix * Revised Preprocessing & Supported Random Forest * Revised to support three models with feature * Further revised prompts * Slight Revision * docs: update contributors (#230) * Revised to support three models with feature * Further revised prompts * Slight Revision * feat: kaggle model and feature (#238) * update first version code * make hypothesis_gen and experiment_builder fit for both feature and model * feat: continue kaggle feature and model coder (#239) * use qlib docker to run qlib models * feature coder ready * model coder ready * fix CI * finish the first round of runner (#240) * Optimized the factor scenario and added the front-end. * fix a small bug * fix a typo * update the kaggle scenario * delete model_template folder * use experiment to run data preprocess script * add source data to scenarios * minor fix * minor bug fix * train.py debug * fixed a bug in train.py and added some TODOs * For Debugging * fix two small bugs in based_exp * fix some bugs * update preprocess * fix a bug in preprocess * fix a bug in train.py * reformat * Follow-up * fix a bug in train.py * fix a bug in workspace * fix a bug in feature duplication * fix a bug in feedback * fix a bug in preprocessed data * fix a bug om feature engineering * fix a ci error * Debugged & Connected * Fixed error on feedback & added other fixes * fix CI errors * fix a CI bug * fix: fix_dotenv_error (#257) * fix_dotenv_error * format with isort * Update rdagent/app/cli.py --------- Co-authored-by: you-n-g <you-n-g@users.noreply.github.com> * chore(main): release 0.2.1 (#249) Release-As: 0.2.1 * init a scenario for kaggle feature engineering * delete error codes * Delete rdagent/app/kaggle_feature/conf.py --------- Co-authored-by: Young <afe.young@gmail.com> Co-authored-by: Taozhi Wang <taozhi.mark.wang@gmail.com> Co-authored-by: you-n-g <you-n-g@users.noreply.github.com> Co-authored-by: cyncyw <47289405+taozhiwang@users.noreply.github.com> Co-authored-by: Xisen-Wang <xisen_application@163.com> Co-authored-by: Haotian Chen <113661982+Hytn@users.noreply.github.com> Co-authored-by: WinstonLiye <1957922024@qq.com> Co-authored-by: WinstonLiyt <104308117+WinstonLiyt@users.noreply.github.com> Co-authored-by: Linlang <30293408+SunsetWolf@users.noreply.github.com>
45 lines
1.4 KiB
Plaintext
45 lines
1.4 KiB
Plaintext
# MODEL_TYPE = "Tabular"
|
|
# BATCH_SIZE = 32
|
|
# NUM_FEATURES = 10
|
|
# NUM_TIMESTEPS = 4
|
|
# NUM_EDGES = 20
|
|
# INPUT_VALUE = 1.0
|
|
# PARAM_INIT_VALUE = 1.0
|
|
|
|
import pickle
|
|
|
|
import torch
|
|
from model import model_cls
|
|
|
|
if MODEL_TYPE == "Tabular":
|
|
input_shape = (BATCH_SIZE, NUM_FEATURES)
|
|
m = model_cls(num_features=input_shape[1])
|
|
data = torch.full(input_shape, INPUT_VALUE)
|
|
elif MODEL_TYPE == "TimeSeries":
|
|
input_shape = (BATCH_SIZE, NUM_FEATURES, NUM_TIMESTEPS)
|
|
m = model_cls(num_features=input_shape[1], num_timesteps=input_shape[2])
|
|
data = torch.full(input_shape, INPUT_VALUE)
|
|
elif MODEL_TYPE == "Graph":
|
|
node_feature = torch.randn(BATCH_SIZE, NUM_FEATURES)
|
|
edge_index = torch.randint(0, BATCH_SIZE, (2, NUM_EDGES))
|
|
m = model_cls(num_features=NUM_FEATURES)
|
|
data = (node_feature, edge_index)
|
|
else:
|
|
raise ValueError(f"Unsupported model type: {MODEL_TYPE}")
|
|
|
|
# Initialize all parameters of `m` to `param_init_value`
|
|
for _, param in m.named_parameters():
|
|
param.data.fill_(PARAM_INIT_VALUE)
|
|
|
|
# Execute the model
|
|
if MODEL_TYPE == "Graph":
|
|
out = m(*data)
|
|
else:
|
|
out = m(data)
|
|
|
|
execution_model_output = out.cpu().detach()
|
|
execution_feedback_str = f"Execution successful, output tensor shape: {execution_model_output.shape}"
|
|
|
|
pickle.dump(execution_model_output, open("execution_model_output.pkl", "wb"))
|
|
pickle.dump(execution_feedback_str, open("execution_feedback_str.pkl", "wb"))
|