Browse Source

[MNT] rewrite database ops to support postgres

tags/v0.3.2
zouxiaochuan 3 years ago
parent
commit
4d16553900
5 changed files with 169 additions and 113 deletions
  1. +9
    -1
      learnware/config.py
  2. +147
    -85
      learnware/market/database_ops.py
  3. +11
    -8
      learnware/market/easy.py
  4. +2
    -1
      learnware/specification/rkme.py
  5. +0
    -18
      requirements.txt

+ 9
- 1
learnware/config.py View File

@@ -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,
}



+ 147
- 85
learnware/market/database_ops.py View File

@@ -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

+ 11
- 8
learnware/market/easy.py View File

@@ -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



+ 2
- 1
learnware/specification/rkme.py View File

@@ -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"]))


+ 0
- 18
requirements.txt View File

@@ -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

Loading…
Cancel
Save