157 lines
6.1 KiB
Python
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)
|