Files
temperature-based-fertility…/code/new_realtime/dataset_wrapper.py
T
Alex Blank 50cf43b9fe added code
2025-05-19 11:11:04 +02:00

157 lines
6.1 KiB
Python

import sys
import argparse
from concurrent.futures import ProcessPoolExecutor
from functools import partial
import logging
import lmdb
import dotenv
from tqdm import tqdm
dotenv.load_dotenv()
from vsm_datascience_common.cycle_database_connection.cycle_data import get_cycle_by_id
from vsm_datascience_common.cycle_database_connection.db_utils import get_cycles_collection
from utils.dataset_creation import get_features, train_scalers, save_scalers, scale_item
from utils.lmdb_utils import save_to_lmdb, load_from_lmdb
from utils.utils import get_variable_from_module
MAX_LMDB_SIZE_IN_MB = 200_000
logger = logging.getLogger(__name__)
logger.setLevel(logging.INFO)
logger.addHandler(logging.StreamHandler(sys.stdout))
def process_id_wrapper(cycle_id: str, feature_config: dict, env: lmdb.Environment):
try:
cycle = get_cycle_by_id(cycle_id)
features = get_features(cycle, feature_config)
save_to_lmdb(env, key=str(cycle_id), dataset=features)
except:
pass
def scaling_wrapper(key_batch: str,
scalers: dict,
max_lmdb_size_in_mb: int,
lmdb_env_dir: str):
env = lmdb.open(lmdb_env_dir, readonly=True, lock=False)
scaled_items = list()
for key in key_batch:
item = load_from_lmdb(env, key)
scaled_item = scale_item(item, scalers)
scaled_items.append(scaled_item)
env.close()
env = lmdb.open(lmdb_env_dir, map_size=max_lmdb_size_in_mb * 1024 * 1024)
for key, scaled_item in zip(key_batch, scaled_items):
save_to_lmdb(env, key=key, dataset=scaled_item)
env.close()
def create_dataset(model_configuration: dict,
lmdb_root_dir: str,
max_lmdb_size_in_mb: int,
max_workers: int) -> None:
feature_config = model_configuration["feature_config"]
logger.info(f"Creating dataset for feature set {feature_config['feature_set_name']}")
# fetch valid cycle ids from database
valid_cycle_ids = [x["_id"] for x in get_cycles_collection().aggregate(
feature_config["filter_criteria_pipeline"] + [
{
"$project": {
"_id": 1
}
}
]
)]
logger.info(f"Fetched {len(valid_cycle_ids)} valid cycle ids from database")
env = lmdb.open(f"{lmdb_root_dir}/{feature_config['feature_set_name']}", map_size=max_lmdb_size_in_mb * 1024 * 1024)
# create features for items
logger.info(f"Creating features for {len(valid_cycle_ids)} cycles")
with ProcessPoolExecutor(max_workers=max_workers) as executor:
list(tqdm(executor.map(partial(process_id_wrapper, feature_config=feature_config, env=env), valid_cycle_ids),
total=len(valid_cycle_ids)))
# convert ids (here bson objectids) to keys for use in lmdb
keys = [str(cycle_id) for cycle_id in valid_cycle_ids]
# train scalers for featues
logger.info("Training scalers for features")
all_scalers = dict()
sample = load_from_lmdb(env, keys[0])
for feature_set in tqdm(feature_config["feature_sets"]):
for feature in feature_config[feature_set]:
print(f"Training scalers for feature {feature['name']} in feature set {feature_set}")
scaler = train_scalers(feature_set, feature["name"], feature["scaler"], sample, env)
if feature_set not in all_scalers:
all_scalers[feature_set] = dict()
all_scalers[feature_set] = all_scalers[feature_set] | scaler
# save scalers
logger.info("Saving scalers to disk")
scaler_dir = f"{lmdb_root_dir}/{feature_config['feature_set_name']}/scalers"
save_scalers(all_scalers, scaler_dir)
# creat scaling batches for less burden on lmdb
batch_size = 1_000
key_batches = [keys[i:i + batch_size] for i in range(0, len(keys), batch_size)]
# scale items
logger.info("Scaling items")
with ProcessPoolExecutor(max_workers=max_workers) as executor:
list(tqdm(executor.map(partial(scaling_wrapper, scalers=all_scalers,
max_lmdb_size_in_mb=max_lmdb_size_in_mb,
lmdb_env_dir=f"{lmdb_root_dir}/{feature_config['feature_set_name']}"),
key_batches),
total=len(key_batches)))
logger.info(f"Dataset creation finished. LMDB saved in {lmdb_root_dir}/{feature_config['feature_set_name']}")
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Dataset Creation Wrapper")
parser.add_argument("--model_config_module",
type=str,
required=True,
help="Path to the model config module")
parser.add_argument("--model_config_variable",
type=str,
required=False,
default="model_configuration",
help="Name of the model config variable, default is 'model_configuration'")
parser.add_argument("--lmdb_dir",
type=str,
required=False,
default="./lmdb_datasets",
help="Path to the lmdb directory")
parser.add_argument("--lmdb_size",
type=int,
required=False,
default=MAX_LMDB_SIZE_IN_MB,
help="Size of the lmdb in MB, default is 200_000")
parser.add_argument("--max_workers",
type=int,
required=False,
default=None,
help="Number of workers for multiprocessing, default is None (use all available cores)")
args = parser.parse_args()
# import and load the model config
model_configuration = get_variable_from_module(
module_path=args.model_config_module,
variable_name=args.model_config_variable
)
create_dataset(model_configuration,
lmdb_root_dir=args.lmdb_dir,
max_lmdb_size_in_mb=args.lmdb_size,
max_workers=args.max_workers)