From 0d209ed92dbb37e985a8f3cba92fc84ab0a907c5 Mon Sep 17 00:00:00 2001 From: chenzx Date: Fri, 21 Apr 2023 16:18:07 +0800 Subject: [PATCH] [MNT] Update image example --- examples/example_image/main.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/examples/example_image/main.py b/examples/example_image/main.py index 6b7a878..e96094a 100644 --- a/examples/example_image/main.py +++ b/examples/example_image/main.py @@ -4,7 +4,7 @@ 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, EnsembleReuser +from learnware.learnware import Learnware, JobSelectorReuser, AveragingReuser import time from learnware.market import EasyMarket, BaseUserInfo @@ -157,7 +157,6 @@ def test_search(gamma=0.1, load_market=True): sorted_score_list, single_learnware_list, mixture_score, mixture_learnware_list = image_market.search_learnware( user_info ) - print(sorted_score_list) l = len(sorted_score_list) acc_list = [] for idx in range(l): @@ -176,7 +175,7 @@ def test_search(gamma=0.1, load_market=True): print(f"mixture reuse loss: {reuse_score}\n") """ - reuse_ensemble = EnsembleReuser(learnware_list=mixture_learnware_list, mode="vote") + reuse_ensemble = AveragingReuser(learnware_list=mixture_learnware_list, mode="vote") ensemble_predict_y = reuse_ensemble.predict(user_data=user_data) ensemble_score = eval_prediction(ensemble_predict_y, user_label) ensemble_score_list.append(ensemble_score)