-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain_PA_hyper.m
More file actions
162 lines (129 loc) · 6.21 KB
/
Copy pathmain_PA_hyper.m
File metadata and controls
162 lines (129 loc) · 6.21 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
clc;
clear;
close all;
warning off;
addpath(genpath('.\'));
%data_sets = ["satimage_multiclass","wine","dna"];
%data_sets = ["procancer","endocrinecancer","cnscancer"];
%data_sets=["ovary_can","breast_can","globun_can","brain_can1","lung_can","pomeroy", "nakayam","sing_procancer"];
data_sets=["alon","bur","chin","chowdary","brain_can","gravier","west","sun","ship","Leukemia"];
%data_sets = ["mirflickr"];
data_set = data_sets(7);
save_results = 0;
xlsx_write = 0;
range = ["b1:j14"];
for f = 1:length(data_set)
percent = [0.9 0.8 0.7 0.6 0.5 ];
%percent = [0 0.1 0.2 0.3 0.4 0.5];
% percent = [0.9];
%
for ii = 1:length(percent)
%% Load a multi-label dataset
dataset = data_set(f);
load(strcat(dataset,".mat"));
%% set all 0 to -1 in target
target(target(:,:)==0) = -1;
%===========for multi-class data==================
target = target';
%=============================
% rng(9,'Twister');
% scurr = rng;
%
%% Randomly select part of data
max_num = 7000;
if size(data,1) > max_num
nRows = size(data,1);
nSample = max_num;
rndIDX = randperm(nRows);
index = rndIDX(1:nSample);
data = data(index, :);
target = target(:,index);
end
%% randomisation of the dataset ()
nRows = size(data,1);
nSample = nRows;
rndIDX = randperm(nRows);
index = rndIDX(1:nSample);
data = data(index, :);
target = target(:,index);
%% normalise
normalise = 1;
if normalise==1
data = svdatanorm(data,'ker');
%data = normalize(data,'zscore');
end
%% Perform n-fold cross validation and obtain evaluation results
rng(15,'Twister');
scurr = rng;
num_fold = 5; num_metric = 5; num_method = 1;
indices = crossvalind('Kfold',size(data,1),num_fold);
Results = zeros(num_metric+1,num_fold,num_method);
for i = 1:num_fold
test = (indices == i);
a = i+1;
if a>num_fold
a = 1;
end
vald = (indices == a );
train = ~test & ~vald;
%disp(['Fold ',num2str(i)]);
fprintf('Fold %d ',i);
kernel='rbf';
data_train = data(logical(train+vald),:);
target_train = target(:,logical(train+vald))';
data_test = data(logical(test),:);
target_test = target(:,logical(test))';
per = percent(ii);
%% ======================Proposed Algo=============================================================
par.c1=0.1; par.c2=0.1; par.c3=0.2; par.c4=0.3; par.c5=0.6; par.c6=0.3; par.c7=0.5;%par.c8=0;
[target_train_incomp] = mask_target_entries(target_train, per);
%[Pre_Labels_train,Pre_Labels_test,time] = myfunupdated(data_train,target_train_incomp,data_test,target_test,par);
[Pre_Labels_train,Pre_Labels_test,time,obj1] = myfunupdatedHyper(data_train,target_train_incomp,data_test,target_test,par);
%% =======================Result Evaluation=======================================================
Results(1,i,1) = time;
% [ExactM_train,HamS_train,MacroF1_train,MicroF1_train,AvePre_train] = Evaluation(Pre_Labels_train',target_train');
[ExactM_test,HamS_test,MacroF1_test,MicroF1_test,AvePre_test] = Evaluation(Pre_Labels_test',target_test');
Results(2:end,i,1) = [ExactM_test,HamS_test,MacroF1_test,MicroF1_test,AvePre_test];
end
ignore = []; Results(:,:,ignore) = [];
meanResults = squeeze(mean(Results,2));
stdResults = squeeze(std(Results,0,2) / sqrt(size(Results,2)));
%addpath("D:\research-track\table generators latex");
meanResults = three_decimals(meanResults);
stdResults = three_decimals(stdResults);
combined = (strcat(arrayfun(@num2str,meanResults,'un',0),'±',arrayfun(@num2str,stdResults,'un',0)));
%% Save the evaluation results
if save_results == 1
filename=strcat("SR_ML_time_Results_",dataset,'.mat');
save(filename,'meanResults','stdResults','-mat');
end
%% Show the experimental results
disp(dataset); disp(per);
for i=1:size(meanResults)
fprintf(' %.2f ± %.2f\n', meanResults(i),stdResults(i));
% disp([meanResults]);
% disp('=========================standard deviation==================================');
% disp(stdResults);
end
%d = {char(dataset),' ',' ',' ',' ',' ',' ', ' '};%,' ',' ',' ',' '};
if xlsx_write == 1
n = {dataset,'percentage of',' missing labels is ',per,'','','','',''};
headers = {per,'algo1','algo2','algo3','algo4','algo5','algo6','algo7','algo8'}; %comparing algo names
%columns = {'Time';'Exact match Train';'Exact match Test';'Hamming Loss Train';'Hamming Loss Test';'Macro F1 Train';'Macro F1 Test';...
% 'Micro F1 Train';'Micro F1 Test';'Avg Precision Train';'Avg Precision Test';};
columns = {'Time';'Exact match';'Hamming Loss';'Macro F1';...
'Micro F1';'Avg Precision';};
%data_as_cell1 = num2cell([meanResults]);
%data_as_cell2 = num2cell([stdResults]);
data_as_cell = combined;
%info_to_write = [headers; data_as_cell1;n;data_as_cell2];
info_to_write = [data_as_cell;];
info_to_write = [columns info_to_write];
info_to_write = [n; info_to_write];
% xlswrite(strcat(dataset,".xls"), info_to_write,range(f));
write_data = cell2table(info_to_write,"VariableNames",[" ","algo1","algo2","algo3","algo4","algo5","algo6","algo7","algo8"]);
filename=strcat(dataset,num2str(per),'.xlsx');
writetable(write_data,filename);
end
end
end