Browse Source

[FIX, ENH] async add_learnware, fix heter typo

tags/v0.3.2
liuht 2 years ago
parent
commit
e7b7e869e2
13 changed files with 74 additions and 241 deletions
  1. +1
    -2
      .gitignore
  2. +1
    -1
      learnware/market/__init__.py
  3. +0
    -177
      learnware/market/hetergeneous/database_ops.py
  4. +0
    -0
      learnware/market/heterogeneous/__init__.py
  5. +0
    -0
      learnware/market/heterogeneous/organizer.py
  6. +41
    -44
      learnware/market/heterogeneous/organizer/__init__.py
  7. +0
    -0
      learnware/market/heterogeneous/organizer/hetero_mapping/__init__.py
  8. +0
    -0
      learnware/market/heterogeneous/organizer/hetero_mapping/feature_extractor.py
  9. +0
    -0
      learnware/market/heterogeneous/organizer/hetero_mapping/trainer.py
  10. +25
    -12
      learnware/market/heterogeneous/searcher.py
  11. +1
    -1
      learnware/market/module.py
  12. +2
    -2
      learnware/specification/system/heter_table.py
  13. +3
    -2
      tests/test_market/test_hetero_market/test_hetero.py

+ 1
- 2
.gitignore View File

@@ -43,5 +43,4 @@ cache/
tmp/
learnware_pool/
PFS/
data/
learnware/market/hetergeneous/.learnware/*
data/

+ 1
- 1
learnware/market/__init__.py View File

@@ -3,7 +3,7 @@ from .base import BaseUserInfo, LearnwareMarket, BaseChecker, BaseOrganizer, Bas
from .evolve_anchor import EvolvedAnchoredOrganizer
from .evolve import EvolvedOrganizer
from .easy import EasyOrganizer, EasySearcher, EasySemanticChecker, EasyStatChecker
from .hetergeneous import HeteroMapTableOrganizer, HeteroSearcher
from .heterogeneous import HeteroMapTableOrganizer, HeteroSearcher

from .classes import CondaChecker
from .module import instantiate_learnware_market

+ 0
- 177
learnware/market/hetergeneous/database_ops.py View File

@@ -1,177 +0,0 @@
import json
import os

from learnware.learnware import get_learnware_from_dirpath
from learnware.logger import get_module_logger
from sqlalchemy import (Column, DateTime, Integer, String, Text, create_engine,
text)
from sqlalchemy.ext.declarative import declarative_base

logger = get_module_logger("database")
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] + "/postgres"
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))
conn.commit()
pass
pass

def update_learnware_semantic_specification(self, id: str, semantic_spec: dict):
with self.engine.connect() as conn:
semantic_spec_str = json.dumps(semantic_spec)
r = conn.execute(
text("UPDATE tb_learnware SET semantic_spec=:semantic_spec WHERE id=:id;"),
dict(id=id, semantic_spec=semantic_spec_str),
)
conn.commit()
pass
pass

def update_learnware_use_flag(self, id: str, use_flag: str):
with self.engine.connect() as conn:
r = conn.execute(
text("UPDATE tb_learnware SET use_flag=:use_flag WHERE id=:id;"),
dict(id=id, use_flag=use_flag),
)
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 = {}
use_flags = {}
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
)
logger.info(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
use_flags[id] = use_flag
max_count = max(max_count, int(id))
pass

return learnware_list, zip_list, folder_list, use_flags, max_count + 1
pass

pass

learnware/market/hetergeneous/__init__.py → learnware/market/heterogeneous/__init__.py View File


learnware/market/hetergeneous/organizer.py → learnware/market/heterogeneous/organizer.py View File


learnware/market/hetergeneous/organizer/__init__.py → learnware/market/heterogeneous/organizer/__init__.py View File

@@ -17,9 +17,9 @@ from ....learnware import Learnware, get_learnware_from_dirpath
from ....logger import get_module_logger
from ....specification.system import HeteroSpecification
from ...base import BaseChecker, BaseUserInfo
from ...easy2 import EasyOrganizer
from ..database_ops import DatabaseOperations
from .config import C as conf
from ...easy import EasyOrganizer
from ...easy.database_ops import DatabaseOperations
from ....config import C as conf
from .hetero_mapping import HeteroMapping, Trainer

logger = get_module_logger("hetero_market")
@@ -27,12 +27,12 @@ logger = get_module_logger("hetero_market")

class HeteroMapTableOrganizer(EasyOrganizer):
def reload_market(self, rebuild=False, auto_update_limit=100):
self.market_store_path = os.path.join(conf.hetero_root_path, self.market_id)
self.market_mapping_path = os.path.join(self.market_store_path, conf.market_model_path)
self.market_store_path = os.path.join(conf.root_path, self.market_id)
self.market_mapping_path = os.path.join(self.market_store_path, "model.bin")
self.learnware_pool_path = os.path.join(self.market_store_path, "learnware_pool")
self.learnware_zip_pool_path = os.path.join(self.market_store_path, "zips")
self.learnware_folder_pool_path = os.path.join(self.market_store_path, "unzipped_learnwares")
self.hetero_mappings_path = os.path.join(self.market_store_path, conf.heter_mapping_path)
self.hetero_mappings_path = os.path.join(self.market_store_path, "hetero_mappings")
self.learnware_list = {} # id:learnware
self.learnware_zip_list = {}
self.learnware_folder_list = {}
@@ -41,8 +41,6 @@ class HeteroMapTableOrganizer(EasyOrganizer):
self.dbops = DatabaseOperations(conf.database_url, "market_" + self.market_id)
self.auto_update = False
self.auto_update_limit = auto_update_limit
self.auto_update_lock = mp.Lock()
self.is_training_in_progress = mp.Value('i', 0)

if rebuild:
logger.warning("Warning! You are trying to clear current database!")
@@ -75,7 +73,6 @@ class HeteroMapTableOrganizer(EasyOrganizer):
self.market_mapping = HeteroMapping()

def reset(self, market_id=None, auto_update=False, auto_update_limit=None, **kwargs):
# model training arguments(model architecture + optimization) set via self.reset
self.auto_update = auto_update
self.market_id = market_id
self.training_args = kwargs
@@ -126,42 +123,45 @@ class HeteroMapTableOrganizer(EasyOrganizer):
self.learnware_zip_list[learnware_id] = target_zip_dir
self.learnware_folder_list[learnware_id] = target_folder_dir
self.use_flags[learnware_id] = learnwere_status
self.count += 1

with self.auto_update_lock:
if self.auto_update and not self.is_training_in_progress.value and self.count - self.last_trained_learnware_num >= self.auto_update_limit:
self.is_training_in_progress.value = 1
curr_learnware_list = copy.deepcopy(self.learnware_list)
train_process = mp.Process(target=self.train, args=(curr_learnware_list.values(),))
train_process.start()
# train_process.join()
self.count += 1

if self.auto_update and self.count - self.last_trained_learnware_num == self.auto_update_limit + 1:
logger.warning(f"Leanwares for training: {self.get_learnware_ids()}")

updated_market_mapping = self.train(
learnware_list=self.learnware_list.values(),
save_dir=self.market_store_path,
**self.training_args
)
logger.warning(f"Market mapping train completed. Now update HeteroSpecification for {self.get_learnware_ids()}")
self.market_mapping = updated_market_mapping
self._update_learnware_list(self.learnware_list.values())
self.last_trained_learnware_num = self.count
return learnware_id, learnwere_status

def train(self, learnware_list: List[Learnware] = None):
learnware_list = learnware_list or self.learnware_list.values()
logger.warning(f"Leanwares for training: {[learnware.id for learnware in learnware_list]}")
allset = self._learnwares_to_dataframes(learnware_list)
self.market_mapping = HeteroMapping(**self.training_args)
@staticmethod
def train(learnware_list: List[Learnware] = None, save_dir: str = None, **kwargs):
allset = HeteroMapTableOrganizer._learnwares_to_dataframes(learnware_list)
market_mapping = HeteroMapping(**kwargs)
market_mapping_trainer = Trainer(
model=self.market_mapping,
model=market_mapping,
train_set_list=allset,
collate_fn=self.market_mapping.collate_fn,
**self.training_args,
collate_fn=market_mapping.collate_fn,
**kwargs,
)
market_mapping_trainer.train()

# auto save whenever market model retrained
market_mapping_trainer.save_model(output_dir=self.market_store_path)

# essential hetero-mapping update for each market learnware when market model retrained
self._update_learnware_list(learnware_list)
self.last_trained_learnware_num = self.count
market_mapping_trainer.train()
market_mapping_trainer.save_model(output_dir=save_dir)

logger.warning(f"Updataed Specification For: {[learnware.id for learnware in learnware_list]}")
return market_mapping

with self.auto_update_lock:
self.is_training_in_progress.value = 0
############################################
# save_model & generateing new specification
# should be moved out of train thread
############################################

def _update_learnware_list(self, learnware_list: List[Learnware]):
try:
@@ -169,7 +169,7 @@ class HeteroMapTableOrganizer(EasyOrganizer):
hetero_spec_path = os.path.join(self.hetero_mappings_path, f"{learnware.id}.npy")
self._update_learnware_specification(learnware, save_path=hetero_spec_path)
except Exception as err:
logger.warning(f"Update learnware HeteroSpecification failed! Due to {err}")
logger.warning(f"Update HeteroSpecification failed! Due to {err}")

def _update_learnware_specification(self, learnware: Learnware, save_path: str) -> Learnware:
specification = learnware.specification
@@ -178,16 +178,16 @@ class HeteroMapTableOrganizer(EasyOrganizer):
learnware_hetero_spec = self.market_mapping.hetero_mapping(learnware_rkme, learnware_features)
learnware.update_stat_spec("HeteroSpecification", learnware_hetero_spec)

# custom hetero spec save path?
learnware_hetero_spec.save(save_path)

def generate_hetero_map_spec(self, user_info: BaseUserInfo) -> HeteroSpecification:
user_rkme = user_info.stat_info["RKMETableSpecification"]
user_features = user_info.semantic_spec["Input"]["Description"].values()
user_features = user_info.get_semantic_spec()["Input"]["Description"].values()
user_hetero_spec = self.market_mapping.hetero_mapping(user_rkme, user_features)
return user_hetero_spec

def _learnwares_to_dataframes(self, learnware_list: List[Learnware]) -> List[pd.DataFrame]:
@staticmethod
def _learnwares_to_dataframes(learnware_list: List[Learnware]) -> List[pd.DataFrame]:
learnware_df_dict = defaultdict(list)
for learnware in learnware_list:
specification = learnware.get_specification()
@@ -198,7 +198,4 @@ class HeteroMapTableOrganizer(EasyOrganizer):
learnware_df_dict[tuple(sorted(learnware_features))].append(learnware_df)

merged_dfs = [pd.concat(dfs) for dfs in learnware_df_dict.values()]
return merged_dfs

def save(self, save_path):
return NotImplementedError
return merged_dfs

learnware/market/hetergeneous/organizer/hetero_mapping/__init__.py → learnware/market/heterogeneous/organizer/hetero_mapping/__init__.py View File


learnware/market/hetergeneous/organizer/hetero_mapping/feature_extractor.py → learnware/market/heterogeneous/organizer/hetero_mapping/feature_extractor.py View File


learnware/market/hetergeneous/organizer/hetero_mapping/trainer.py → learnware/market/heterogeneous/organizer/hetero_mapping/trainer.py View File


learnware/market/hetergeneous/searcher.py → learnware/market/heterogeneous/searcher.py View File

@@ -1,10 +1,12 @@
from typing import Tuple, List, Union

import numpy as np

from ...learnware import Learnware
from ...logger import get_module_logger
from ...specification import HeteroSpecification
from ..base import BaseSearcher, BaseUserInfo
from ..easy2 import EasySearcher
from ..easy import EasySearcher
from ..utils import parse_specification_type
from .organizer import HeteroMapTableOrganizer

@@ -38,7 +40,7 @@ class HeteroMapTableSearcher(EasySearcher):
) -> Tuple[List[float], List[Learnware]]:
hetero_spec_list = [learnware.specification.get_stat_spec_by_name("HeteroSpecification") for learnware in learnware_list]
mmd_dist_list = []
for hetero_spec in hetero_spec_list:
for idx, hetero_spec in enumerate(hetero_spec_list):
mmd_dist = hetero_spec.dist(user_hetero_spec)
mmd_dist_list.append(mmd_dist)
@@ -83,15 +85,6 @@ class HeteroMapTableSearcher(EasySearcher):
logger.info(f"After filter by hetero spec, learnware_list length is {len(single_learnware_list)}")
return sorted_score_list, single_learnware_list, None, None

# for learnware in learnware_list:
# learnware_hetero_spec = learnware.specification.get_stat_spec_by_name("HeteroSpecification")
# mmd_dist = learnware_hetero_spec.dist(user_hetero_spec)
# if target_learnware is None or mmd_dist < min_dist:
# min_dist = mmd_dist
# target_learnware = learnware
# return target_learnware

def reset(self, organizer):
self.learnware_oganizer = organizer

@@ -103,6 +96,26 @@ class HeteroSearcher(EasySearcher):
def reset(self, organizer):
super().reset(organizer)
self.hetero_stat_searcher.reset(organizer)
@staticmethod
def check_user_info(user_info: BaseUserInfo):
try:
user_stat_spec = user_info.get_stat_info("RKMETableSpecification")
user_input_shape = user_stat_spec.get_z().shape[1]

user_input_description = user_info.get_semantic_spec()["Input"]

user_description_dim = int(user_input_description["Dimension"])
user_description_feature_num = len(user_input_description["Description"])

if user_input_shape != user_description_dim or user_input_shape != user_description_feature_num:
logger.warning("User data feature dimensions mismatch with semantic specification")
return False
return True
except:
logger.info(f"No heterogeneous search information provided. Use homogeneous search instead.")
return False

def __call__(
self, user_info: BaseUserInfo, check_status: int = None, max_search_num: int = 5, search_method: str = "greedy"
@@ -114,7 +127,7 @@ class HeteroSearcher(EasySearcher):
return [], [], 0.0, []

if parse_specification_type(stat_specs=user_info.stat_info) is not None:
if "Input" in user_info.semantic_spec and user_info.semantic_spec["Input"]["Description"] is not None:
if self.check_user_info(user_info):
return self.hetero_stat_searcher(learnware_list, user_info)
else:
return self.stat_searcher(learnware_list, user_info, max_search_num, search_method)

+ 1
- 1
learnware/market/module.py View File

@@ -1,6 +1,6 @@
from .base import LearnwareMarket
from .easy import EasyOrganizer, EasySearcher, EasySemanticChecker, EasyStatChecker
from .hetergeneous import HeteroMapTableOrganizer, HeteroSearcher
from .heterogeneous import HeteroMapTableOrganizer, HeteroSearcher

MARKET_CONFIG = {
"easy": {


+ 2
- 2
learnware/specification/system/heter_table.py View File

@@ -29,10 +29,10 @@ class HeteroSpecification(SystemStatsSpecification):
super(HeteroSpecification, self).__init__(type=self.__class__.__name__)

def get_z(self) -> np.ndarray:
return self.z.detach().cpu().numpy
return self.z.detach().cpu().numpy()

def get_beta(self) -> np.ndarray:
return self.beta.detach().cpu().numpy
return self.beta.detach().cpu().numpy()

def generate_stat_spec_from_system(self, heter_embedding: np.ndarray, rkme_spec: RKMETableSpecification):
self.beta = rkme_spec.beta.to(self.device)


+ 3
- 2
tests/test_market/test_hetero_market/test_hetero.py View File

@@ -155,6 +155,7 @@ class TestMarket(unittest.TestCase):
hetero_market = self._init_learnware_market()
self.test_prepare_learnware_randomly(learnware_num)
self.learnware_num = learnware_num
hetero_market.learnware_organizer.reset(auto_update=True, auto_update_limit=learnware_num)

print("Total Item:", len(hetero_market))
assert len(hetero_market) == 0, f"The market should be empty!"
@@ -173,8 +174,8 @@ class TestMarket(unittest.TestCase):
print("Available ids After Uploading Learnwares:", curr_inds)
assert len(curr_inds) == self.learnware_num, f"The number of learnwares must be {self.learnware_num}!"

organizer=hetero_market.learnware_organizer
organizer.train()
# organizer=hetero_market.learnware_organizer
# organizer.train(hetero_market.learnware_organizer.learnware_list.values())
return hetero_market

def test_search_semantics(self, learnware_num=5):


Loading…
Cancel
Save