Browse Source

Merge branch 'dev' of git.nju.edu.cn:learnware/learnware-market into dev

tags/v0.3.2
Gene 3 years ago
parent
commit
5ef9ef4309
2 changed files with 20 additions and 9 deletions
  1. +19
    -8
      examples/example_image/main.py
  2. +1
    -1
      learnware/learnware/reuse.py

+ 19
- 8
examples/example_image/main.py View File

@@ -1,9 +1,10 @@
import numpy as np
import torch
import get_data
from get_data import *
import os
import random
from utils import generate_uploader, generate_user, ImageDataLoader, train, eval_prediction
from learnware.learnware import Learnware, JobSelectorReuser
import time

from learnware.market import EasyMarket, BaseUserInfo
@@ -58,9 +59,9 @@ user_senmantic = {

def prepare_data():
if dataset == "cifar10":
X_train, y_train, X_test, y_test = get_data.get_cifar10(data_root)
X_train, y_train, X_test, y_test = get_cifar10(data_root)
elif dataset == "mnist":
X_train, y_train, X_test, y_test = get_data.get_mnist(data_root)
X_train, y_train, X_test, y_test = get_mnist(data_root)
else:
return
generate_uploader(X_train, y_train, n_uploaders=n_uploaders, data_save_root=uploader_save_root)
@@ -130,7 +131,7 @@ def prepare_market():
logger.info("Available ids: " + str(curr_inds))


def test_search(load_market=True):
def test_search(gamma=0.1, load_market=True):
if load_market:
image_market = EasyMarket()
else:
@@ -141,12 +142,13 @@ def test_search(load_market=True):
select_list = []
avg_list = []
improve_list = []
job_selector_score_list = []
for i in range(n_users):
user_data_path = os.path.join(user_save_root, "user_%d_X.npy" % (i))
user_label_path = os.path.join(user_save_root, "user_%d_y.npy" % (i))
user_data = np.load(user_data_path)
user_label = np.load(user_label_path)
user_stat_spec = specification.utils.generate_rkme_spec(X=user_data, gamma=0.1, cuda_idx=0)
user_stat_spec = specification.utils.generate_rkme_spec(X=user_data, gamma=gamma, cuda_idx=0)
user_info = BaseUserInfo(
id=f"user_{i}", semantic_spec=user_senmantic, stat_info={"RKMEStatSpecification": user_stat_spec}
)
@@ -163,17 +165,26 @@ def test_search(load_market=True):
acc = eval_prediction(pred_y, user_label)
acc_list.append(acc)
logger.info("search rank: %d, score: %.3f, learnware_id: %s, acc: %.3f" % (idx, score, learnware.id, acc))

# test reuse
"""
reuse_baseline = JobSelectorReuser(learnware_list=mixture_learnware_list)
reuse_predict = reuse_baseline.predict(user_data=user_data)
reuse_score = eval_prediction(reuse_predict, user_label)
job_selector_score_list.append(reuse_score)
"""
print(f"mixture reuse loss: {reuse_score}\n")
select_list.append(acc_list[0])
avg_list.append(np.mean(acc_list))
improve_list.append((acc_list[0] - np.mean(acc_list)) / np.mean(acc_list))
logger.info(
"Accuracy of selected learnware: %.3f, Average performance: %.3f" % (np.mean(select_list), np.mean(avg_list))
"Accuracy of selected learnware: %.3f +/- %.3f, Average performance: %.3f +/- %.3f"
% (np.mean(select_list), np.std(select_list), np.mean(avg_list), np.std(avg_list))
)
logger.info("Average performance improvement: %.3f" % (np.mean(improve_list)))
# logger.info("Average Job Selector Reuse Performance: %.3f +/- %.3f"%(np.mean(job_selector_score_list), np.std(job_selector_score_list)))


if __name__ == "__main__":
# prepare_data()
# prepare_model()
test_search()
test_search(load_market=True)

+ 1
- 1
learnware/learnware/reuse.py View File

@@ -200,7 +200,7 @@ class JobSelectorReuser(BaseReuser):
boosting_type="gbdt",
seed=0,
)
train_y = train_y.astype(np.int)
train_y = train_y.astype(int)
model.fit(train_x, train_y, eval_set=[(val_x, val_y)], verbose=-1, early_stopping_rounds=300)
pred_y = model.predict(org_train_x)
score = accuracy_score(pred_y, org_train_y)


Loading…
Cancel
Save