diff --git a/learnware/config.py b/learnware/config.py index 7eaf898..6e51bac 100644 --- a/learnware/config.py +++ b/learnware/config.py @@ -1,6 +1,7 @@ import os import copy import logging +import json class Config: @@ -8,6 +9,13 @@ class Config: self.__dict__["_default_config"] = copy.deepcopy(default_conf) # avoiding conflictions with __getattr__ self.reset() + config_file = os.path.join(self.root_path, "config.json") + if os.path.exists(config_file): + with open(config_file, "r") as f: + self.__dict__["_config"].update(json.load(f)) + pass + pass + def __getitem__(self, key): return self.__dict__["_config"][key] @@ -130,7 +138,7 @@ _DEFAULT_CONFIG = { "yaml_file": "learnware.yaml", "module_file": "__init__.py", }, - "database_path": DATABASE_PATH, + "database_url": f"sqlite:///{DATABASE_PATH}", "max_reduced_set_size": 1310720, } diff --git a/learnware/market/database_ops.py b/learnware/market/database_ops.py index e0c52ee..0a8e6f5 100644 --- a/learnware/market/database_ops.py +++ b/learnware/market/database_ops.py @@ -1,89 +1,151 @@ +from sqlalchemy.ext.declarative import declarative_base +from sqlalchemy import create_engine, text +from sqlalchemy import ( + Column, Integer, Text, DateTime, String +) import os import json -import sqlite3 -from copy import deepcopy - -from ..logger import get_module_logger from ..learnware import get_learnware_from_dirpath -from ..config import C - -logger = get_module_logger("database_ops") - - -def init_empty_db(func): - def wrapper(market_id, *args, **kwargs): - conn = sqlite3.connect(os.path.join(C.database_path, f"market_{market_id}.db")) - cur = conn.cursor() - listOfTables = cur.execute( - """SELECT name FROM sqlite_master WHERE type='table' AND name='LEARNWARE'; """ - ).fetchall() - if len(listOfTables) == 0: - logger.info("Initializing Database in %s..." % (os.path.join(C.database_path, f"market_{market_id}.db"))) - cur.execute( - """CREATE TABLE LEARNWARE - (ID CHAR(10) PRIMARY KEY NOT NULL, - SEMANTIC_SPEC TEXT NOT NULL, - ZIP_PATH TEXT NOT NULL, - FOLDER_PATH TEXT NOT NULL, - USE_FLAG TEXT NOT NULL);""" + + +DeclarativeBase = declarative_base() + + +class Learnware(DeclarativeBase): + __tablename__ = 'tb_learnware' + + id = Column(String(10), primary_key=True, nullable=False) + semantic_spec = Column(Text, nullable=False) + zip_path = Column(Text, nullable=False) + folder_path = Column(Text, nullable=False) + use_flag = Column(Text, nullable=False) + + pass + + +class DatabaseOperations(object): + + def __init__(self, url: str, database_name: str): + if url.startswith("sqlite"): + url = os.path.join(url, f"{database_name}.db") + else: + url = f"{url}/{database_name}" + pass + + self.url = url + self.create_database_if_not_exists(url) + + pass + + + def create_database_if_not_exists(self, url): + database_exists = True + + if url.startswith("sqlite"): + # it is sqlite + start = url.find(":///") + path = url[start+4:] + if os.path.exists(path): + database_exists = True + pass + else: + database_exists = False + os.makedirs(os.path.dirname(path), exist_ok=True) + pass + pass + elif self.url.startswith("postgresql"): + # it is postgresql + dbname_start = url.rfind("/") + dbname = url[dbname_start+1:] + url_no_dbname = url[:dbname_start] + engine = create_engine(url_no_dbname) + + with engine.connect() as conn: + result = conn.execute(text("SELECT datname FROM pg_database;")) + db_list = set() + + for row in result.fetchall(): + db_list.add(row[0].lower()) + pass + + if dbname.lower() not in db_list: + database_exists = False + conn.execution_options(isolation_level="AUTOCOMMIT").execute( + text("CREATE DATABASE {0};".format(dbname))) + pass + else: + database_exists = True + pass + pass + engine.dispose() + pass + else: + raise Exception(f"Unsupported database url: {self.url}") + pass + + self.engine = create_engine(url, future=True) + + if not database_exists: + DeclarativeBase.metadata.create_all(self.engine) + pass + pass + + def clear_learnware_table(self): + with self.engine.connect() as conn: + conn.execute(text("DELETE FROM tb_learnware;")) + conn.commit() + pass + pass + + def add_learnware(self, id: str, semantic_spec: dict, zip_path, folder_path, use_flag: str): + with self.engine.connect() as conn: + semantic_spec_str = json.dumps(semantic_spec) + conn.execute( + text( + ("INSERT INTO tb_learnware (id, semantic_spec, zip_path, folder_path, use_flag)" + "VALUES (:id, :semantic_spec, :zip_path, :folder_path, :use_flag);") + ), + dict(id=id, semantic_spec=semantic_spec_str, zip_path=zip_path, + folder_path=folder_path, use_flag=use_flag) + ) + conn.commit() + pass + pass + + def delete_learnware(self, id: str): + with self.engine.connect() as conn: + conn.execute( + text("DELETE FROM tb_learnware WHERE id=:id;"), + dict(id=id) ) - logger.info("Database Built!") - kwargs["cur"] = cur - item = func(*args, **kwargs) - conn.commit() - conn.close() - return item - - return wrapper - - -# Clear Learnware Database -# !!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!! -# !!!!! !!!!! -# !!!!! Do NOT use unless highly necessary !!!!! -# !!!!! !!!!! -# !!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!! -@init_empty_db -def clear_learnware_table(cur): - logger.warning("!!! Drop Learnware Table !!!") - cur.execute("DROP TABLE LEARNWARE") - - -@init_empty_db -def add_learnware_to_db(id: str, semantic_spec: dict, zip_path: str, folder_path: str, use_flag: str, cur): - semantic_spec_str = json.dumps(semantic_spec) - cur.execute( - "INSERT INTO LEARNWARE (ID,SEMANTIC_SPEC,ZIP_PATH,FOLDER_PATH,USE_FLAG) \ - VALUES ('%s', '%s', '%s', '%s', '%s')" - % (id, semantic_spec_str, zip_path, folder_path, use_flag) - ) - - -@init_empty_db -def delete_learnware_from_db(id: str, cur): - cur.execute("DELETE from LEARNWARE where ID='%s';" % (id)) - - -@init_empty_db -def load_market_from_db(cur): - logger.info("Reload from Database") - cursor = cur.execute("SELECT id, semantic_spec, zip_path, FOLDER_PATH from LEARNWARE") - - learnware_list = {} - zip_list = {} - folder_list = {} - max_count = 0 - - for id, semantic_spec, zip_path, folder_path in cursor: - semantic_spec_dict = json.loads(semantic_spec) - new_learnware = get_learnware_from_dirpath( - id=id, semantic_spec=semantic_spec_dict, learnware_dirpath=folder_path - ) - - learnware_list[id] = new_learnware - zip_list[id] = zip_path - folder_list[id] = folder_path - max_count = max(max_count, int(id)) - - logger.info("Market Reloaded from DB.") - return learnware_list, zip_list, folder_list, max_count + 1 + conn.commit() + pass + pass + + def load_market(self): + with self.engine.connect() as conn: + cursor = conn.execute(text("SELECT id, semantic_spec, zip_path, folder_path, use_flag FROM tb_learnware;")) + + learnware_list = {} + zip_list = {} + folder_list = {} + max_count = 0 + + for id, semantic_spec, zip_path, folder_path, use_flag in cursor: + id = id.strip() + semantic_spec_dict = json.loads(semantic_spec) + new_learnware = get_learnware_from_dirpath( + id=id, semantic_spec=semantic_spec_dict, learnware_dirpath=folder_path + ) + print(f'load learnware: {id}') + learnware_list[id] = new_learnware + # assert new_learnware is not None + zip_list[id] = zip_path + folder_list[id] = folder_path + max_count = max(max_count, int(id)) + pass + + return learnware_list, zip_list, folder_list, max_count + 1 + pass + + pass \ No newline at end of file diff --git a/learnware/market/easy.py b/learnware/market/easy.py index b69a664..5f504e8 100644 --- a/learnware/market/easy.py +++ b/learnware/market/easy.py @@ -9,7 +9,7 @@ from cvxopt import solvers, matrix from typing import Tuple, Any, List, Union, Dict from .base import BaseMarket, BaseUserInfo -from .database_ops import load_market_from_db, add_learnware_to_db, delete_learnware_from_db, clear_learnware_table +from .database_ops import DatabaseOperations from ..learnware import Learnware, get_learnware_from_dirpath from ..specification import RKMEStatSpecification, Specification @@ -54,6 +54,7 @@ class EasyMarket(BaseMarket): self.learnware_folder_list = {} self.count = 0 self.semantic_spec_list = conf.semantic_specs + self.dbops = DatabaseOperations(conf.database_url, 'market_' + self.market_id) self.reload_market(rebuild=rebuild) # Automatically reload the market logger.info("Market Initialized!") @@ -61,7 +62,7 @@ class EasyMarket(BaseMarket): if rebuild: logger.warning("Warning! You are trying to clear current database!") try: - clear_learnware_table(market_id=self.market_id) + self.dbops.clear_learnware_table() rmtree(self.learnware_pool_path) except: pass @@ -69,9 +70,7 @@ class EasyMarket(BaseMarket): os.makedirs(self.learnware_pool_path, exist_ok=True) os.makedirs(self.learnware_zip_pool_path, exist_ok=True) os.makedirs(self.learnware_folder_pool_path, exist_ok=True) - self.learnware_list, self.learnware_zip_list, self.learnware_folder_list, self.count = load_market_from_db( - market_id=self.market_id - ) + self.learnware_list, self.learnware_zip_list, self.learnware_folder_list, self.count = self.dbops.load_market() @classmethod def check_learnware(cls, learnware: Learnware) -> int: @@ -205,10 +204,10 @@ class EasyMarket(BaseMarket): if new_learnware is None: return None, self.INVALID_LEARNWARE + check_flag = self.check_learnware(new_learnware) - add_learnware_to_db( - market_id=self.market_id, + self.dbops.add_learnware( id=id, semantic_spec=semantic_spec, zip_path=target_zip_dir, @@ -655,6 +654,7 @@ class EasyMarket(BaseMarket): else: user_rkme = user_info.stat_info["RKMEStatSpecification"] learnware_list = self._filter_by_rkme_spec_dimension(learnware_list, user_rkme) + print('after filter by rkme dimension, learnware_list length is %d' % len(learnware_list)) sorted_dist_list, single_learnware_list = self._search_by_rkme_spec_single(learnware_list, user_rkme) if search_method == "auto": @@ -679,10 +679,13 @@ class EasyMarket(BaseMarket): sorted_score_list = merge_score_list[:-1] mixture_score = merge_score_list[-1] + print('after search by rkme spec, learnware_list length is %d' % len(learnware_list)) # filter learnware with low score sorted_score_list, single_learnware_list = self._filter_by_rkme_spec_single( sorted_score_list, single_learnware_list ) + + print('after filter by rkme spec, learnware_list length is %d' % len(learnware_list)) return sorted_score_list, single_learnware_list, mixture_score, mixture_learnware_list def delete_learnware(self, id: str) -> bool: @@ -710,7 +713,7 @@ class EasyMarket(BaseMarket): self.learnware_list.pop(id) self.learnware_zip_list.pop(id) self.learnware_folder_list.pop(id) - delete_learnware_from_db(market_id=self.market_id, id=id) + self.dbops.delete_learnware(id=id) return True diff --git a/learnware/specification/rkme.py b/learnware/specification/rkme.py index cdf84f7..f3bddc5 100644 --- a/learnware/specification/rkme.py +++ b/learnware/specification/rkme.py @@ -388,7 +388,8 @@ class RKMEStatSpecification(BaseStatSpecification): # Load JSON file: load_path = filepath if os.path.exists(load_path): - obj_text = codecs.open(load_path, "r", encoding="utf-8").read() + with codecs.open(load_path, "r", encoding="utf-8") as fin: + obj_text = fin.read() rkme_load = json.loads(obj_text) rkme_load["device"] = choose_device(rkme_load["cuda_idx"]) rkme_load["z"] = torch.from_numpy(np.array(rkme_load["z"])) diff --git a/requirements.txt b/requirements.txt deleted file mode 100644 index 14584a5..0000000 --- a/requirements.txt +++ /dev/null @@ -1,18 +0,0 @@ -cvxopt==1.3.1 -faiss==1.5.3 -faiss_cpu==1.7.4 -fire==0.5.0 -joblib==1.1.0 -lightgbm==3.3.5 -matplotlib==3.5.1 -numpy==1.21.5 -pandas==1.4.2 -psutil==5.8.0 -PyYAML==6.0 -requests==2.27.1 -scikit_learn==1.2.2 -scipy==1.7.3 -setuptools==61.2.0 -torch==2.0.1 -torchvision==0.15.2 -tqdm==4.64.0