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