-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathSlopeOne.java
More file actions
133 lines (99 loc) · 3.52 KB
/
Copy pathSlopeOne.java
File metadata and controls
133 lines (99 loc) · 3.52 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
import java.io.*;
import java.util.*;
public class SlopeOne {
static Map<String, Map<Integer, Double>> data = new LinkedHashMap<>();
public static void main(String[] args) throws Exception {
if (args.length != 1) {
System.out.println("Usage : java SlopeOne evaluations.txt");
return;
}
lireFichier(args[0]);
double mae = calculerMAE();
System.out.printf("MAE Slope One = %.4f%n", mae);
}
static void lireFichier(String fichier) throws Exception {
Evaluations evaluations = new Evaluations(fichier);
for (String user : evaluations.utilisateurs()) {
Map<Integer, Double> notes = new HashMap<>();
for (Map.Entry<String, Double> entry : evaluations.evaluations(user)) {
notes.put(Integer.parseInt(entry.getKey()), entry.getValue());
}
data.put(user, notes);
}
}
static double calculerMAE() {
double sommeErreur = 0.0;
int total = 0;
for (String user : data.keySet()) {
for (Integer item : new ArrayList<>(data.get(user).keySet())) {
double vraieNote = data.get(user).get(item);
data.get(user).remove(item);
double prediction = predire(user, item);
data.get(user).put(item, vraieNote);
sommeErreur += Math.abs(vraieNote - prediction);
total++;
}
}
return sommeErreur / total;
}
static double predire(String user, int itemCible) {
Map<Integer, Double> notesUser = data.get(user);
double numerateur = 0.0;
double denominateur = 0.0;
for (Integer itemConnu : notesUser.keySet()) {
Difference diff = calculerDifference(itemCible, itemConnu);
if (diff.count > 0) {
numerateur += (notesUser.get(itemConnu) + diff.valeur) * diff.count;
denominateur += diff.count;
}
}
if (denominateur == 0) {
return moyenne(notesUser);
}
double prediction = numerateur / denominateur;
return limiter(prediction);
}
static Difference calculerDifference(int itemA, int itemB) {
double somme = 0.0;
int count = 0;
for (String user : data.keySet()) {
Map<Integer, Double> notes = data.get(user);
if (notes.containsKey(itemA) && notes.containsKey(itemB)) {
somme += notes.get(itemA) - notes.get(itemB);
count++;
}
}
if (count == 0) return new Difference(0.0, 0);
return new Difference(somme / count, count);
}
static double moyenne(Map<Integer, Double> notes) {
if (notes.isEmpty()) return moyenneGlobale();
double somme = 0.0;
for (double v : notes.values()) somme += v;
return somme / notes.size();
}
static double moyenneGlobale() {
double somme = 0.0;
int total = 0;
for (Map<Integer, Double> notes : data.values()) {
for (double v : notes.values()) {
somme += v;
total++;
}
}
return somme / total;
}
static double limiter(double note) {
if (note < 1) return 1;
if (note > 8) return 8;
return note;
}
static class Difference {
double valeur;
int count;
Difference(double valeur, int count) {
this.valeur = valeur;
this.count = count;
}
}
}