Le clustering K-means partitionne un jeu de données en k groupes en affinant itérativement les centroïdes des clusters jusqu'à leur stabilisation ou jusqu'à l'atteinte de la limite d'itérations. Cette rubrique explique comment implémenter le clustering K-means à l'aide de l'API MaxCompute Graph (Java).
Fonctionnement
À chaque itération, l'algorithme attribue chaque point de données au centroïde le plus proche, puis recalcule le centroïde comme la moyenne de tous les points du cluster. Il s'arrête lorsque le déplacement de tous les centroïdes est inférieur au seuil de convergence ou lorsque le nombre maximal d'itérations est atteint.
Étapes de l'algorithme
Sélectionnez les centroïdes initiaux pour les k clusters.
Pour chaque point de données, calculez sa distance euclidienne au carré par rapport aux k centroïdes et attribuez-le au cluster le plus proche.
Recalculez chaque centroïde comme la moyenne arithmétique de tous les points attribués à ce cluster.
Si le déplacement de tous les centroïdes est inférieur au seuil de convergence, arrêtez l'algorithme. Sinon, recommencez à l'étape 2.
Paramètres clés
| Paramètre | Valeur par défaut | Description |
|---|---|---|
| Seuil de convergence | 0,05 | Distance euclidienne maximale dont un centroïde peut se déplacer entre deux itérations avant que l'algorithme ne le considère comme convergé. Ce calcul utilise la racine carrée de la somme des différences de coordonnées au carré. |
| Nombre maximal d'itérations | 30 | Limite supérieure du nombre de supersteps. Le job s'arrête même si les centroïdes n'ont pas entièrement convergé. |
| Métrique de distance | Euclidienne au carré | La distance au carré sert à comparer les attributions de cluster. La racine carrée n'est calculée que lors de la vérification de la convergence dans terminate(). |
Implémentation
L'exemple suivant implémente un clustering K-means avec quatre classes, basé sur l'API MaxCompute Graph.
Aperçu des composants
Avant d'examiner le code, passez en revue la responsabilité de chaque composant :
**KmeansVertex** représente un seul point de données sous forme de Tuple (vecteur de caractéristiques). Sa méthode compute() transmet la valeur du sommet à l'agrégateur en appelant context.aggregate().
**KmeansVertexReader** charge la table d'entrée et crée un sommet par enregistrement. Le numéro d'enregistrement devient l'ID du sommet ; toutes les colonnes de l'enregistrement constituent la valeur du sommet sous forme de Tuple.
**KmeansAggrValue** conserve l'état d'agrégation partagé entre les workers, avec trois champs Tuple :
centers: coordonnées actuelles des centroïdessums: sommes accumulées des coordonnées par clustercounts: nombre de points de données par cluster
**KmeansAggregator** encapsule la logique principale de l'algorithme via quatre méthodes :
| Méthode | Responsabilité |
|---|---|
createInitialValue() |
Lors de la superstep 0, lit les centroïdes initiaux à partir du fichier cache nommé "centers". Lors des supersteps suivantes, récupère les centroïdes calculés lors de l'itération précédente via getLastAggregatedValue(0). |
aggregate() |
Pour chaque sommet, trouve le centroïde le plus proche à l'aide de la distance euclidienne au carré et met à jour la sum et le count de ce cluster. |
merge() |
Combine les sums et les counts partiels collectés auprès des workers exécutés en parallèle. |
terminate() |
Calcule les nouveaux centroïdes à partir des sums et des counts. Si le déplacement de tous les centroïdes est inférieur à 0,05 ou si le nombre maximal d'itérations est atteint, écrit les centroïdes finaux dans la table de sortie et renvoie true (arrêt). Sinon, renvoie false (continuation). |
Exemple de code
import java.io.DataInput;
import java.io.DataOutput;
import java.io.IOException;
import org.apache.log4j.Logger;
import com.aliyun.odps.io.WritableRecord;
import com.aliyun.odps.graph.Aggregator;
import com.aliyun.odps.graph.ComputeContext;
import com.aliyun.odps.graph.GraphJob;
import com.aliyun.odps.graph.GraphLoader;
import com.aliyun.odps.graph.MutationContext;
import com.aliyun.odps.graph.Vertex;
import com.aliyun.odps.graph.WorkerContext;
import com.aliyun.odps.io.DoubleWritable;
import com.aliyun.odps.io.LongWritable;
import com.aliyun.odps.io.NullWritable;
import com.aliyun.odps.data.TableInfo;
import com.aliyun.odps.io.Text;
import com.aliyun.odps.io.Tuple;
import com.aliyun.odps.io.Writable;
public class Kmeans {
private final static Logger LOG = Logger.getLogger(Kmeans.class);
public static class KmeansVertex extends
Vertex<Text, Tuple, NullWritable, NullWritable> {
@Override
public void compute(
ComputeContext<Text, Tuple, NullWritable, NullWritable> context,
Iterable<NullWritable> messages) throws IOException {
context.aggregate(getValue());
}
}
public static class KmeansVertexReader extends
GraphLoader<Text, Tuple, NullWritable, NullWritable> {
@Override
public void load(LongWritable recordNum, WritableRecord record,
MutationContext<Text, Tuple, NullWritable, NullWritable> context)
throws IOException {
KmeansVertex vertex = new KmeansVertex();
vertex.setId(new Text(String.valueOf(recordNum.get())));
vertex.setValue(new Tuple(record.getAll()));
context.addVertexRequest(vertex);
}
}
public static class KmeansAggrValue implements Writable {
Tuple centers = new Tuple();
Tuple sums = new Tuple();
Tuple counts = new Tuple();
@Override
public void write(DataOutput out) throws IOException {
centers.write(out);
sums.write(out);
counts.write(out);
}
@Override
public void readFields(DataInput in) throws IOException {
centers = new Tuple();
centers.readFields(in);
sums = new Tuple();
sums.readFields(in);
counts = new Tuple();
counts.readFields(in);
}
@Override
public String toString() {
return "centers " + centers.toString() + ", sums " + sums.toString()
+ ", counts " + counts.toString();
}
}
public static class KmeansAggregator extends Aggregator<KmeansAggrValue> {
@SuppressWarnings("rawtypes")
@Override
public KmeansAggrValue createInitialValue(WorkerContext context)
throws IOException {
KmeansAggrValue aggrVal = null;
if (context.getSuperstep() == 0) {
aggrVal = new KmeansAggrValue();
aggrVal.centers = new Tuple();
aggrVal.sums = new Tuple();
aggrVal.counts = new Tuple();
byte[] centers = context.readCacheFile("centers");
String lines[] = new String(centers).split("\n");
for (int i = 0; i < lines.length; i++) {
String[] ss = lines[i].split(",");
Tuple center = new Tuple();
Tuple sum = new Tuple();
for (int j = 0; j < ss.length; ++j) {
center.append(new DoubleWritable(Double.valueOf(ss[j].trim())));
sum.append(new DoubleWritable(0.0));
}
LongWritable count = new LongWritable(0);
aggrVal.sums.append(sum);
aggrVal.counts.append(count);
aggrVal.centers.append(center);
}
} else {
aggrVal = (KmeansAggrValue) context.getLastAggregatedValue(0);
}
return aggrVal;
}
@Override
public void aggregate(KmeansAggrValue value, Object item) {
int min = 0;
double mindist = Double.MAX_VALUE;
Tuple point = (Tuple) item;
for (int i = 0; i < value.centers.size(); i++) {
Tuple center = (Tuple) value.centers.get(i);
// use Euclidean Distance, no need to calculate sqrt
double dist = 0.0d;
for (int j = 0; j < center.size(); j++) {
double v = ((DoubleWritable) point.get(j)).get()
- ((DoubleWritable) center.get(j)).get();
dist += v * v;
}
if (dist < mindist) {
mindist = dist;
min = i;
}
}
// update sum and count
Tuple sum = (Tuple) value.sums.get(min);
for (int i = 0; i < point.size(); i++) {
DoubleWritable s = (DoubleWritable) sum.get(i);
s.set(s.get() + ((DoubleWritable) point.get(i)).get());
}
LongWritable count = (LongWritable) value.counts.get(min);
count.set(count.get() + 1);
}
@Override
public void merge(KmeansAggrValue value, KmeansAggrValue partial) {
for (int i = 0; i < value.sums.size(); i++) {
Tuple sum = (Tuple) value.sums.get(i);
Tuple that = (Tuple) partial.sums.get(i);
for (int j = 0; j < sum.size(); j++) {
DoubleWritable s = (DoubleWritable) sum.get(j);
s.set(s.get() + ((DoubleWritable) that.get(j)).get());
}
}
for (int i = 0; i < value.counts.size(); i++) {
LongWritable count = (LongWritable) value.counts.get(i);
count.set(count.get() + ((LongWritable) partial.counts.get(i)).get());
}
}
@SuppressWarnings("rawtypes")
@Override
public boolean terminate(WorkerContext context, KmeansAggrValue value)
throws IOException {
// compute new centers
Tuple newCenters = new Tuple(value.sums.size());
for (int i = 0; i < value.sums.size(); i++) {
Tuple sum = (Tuple) value.sums.get(i);
Tuple newCenter = new Tuple(sum.size());
LongWritable c = (LongWritable) value.counts.get(i);
for (int j = 0; j < sum.size(); j++) {
DoubleWritable s = (DoubleWritable) sum.get(j);
double val = s.get() / c.get();
newCenter.set(j, new DoubleWritable(val));
// reset sum for next iteration
s.set(0.0d);
}
// reset count for next iteration
c.set(0);
newCenters.set(i, newCenter);
}
// update centers
Tuple oldCenters = value.centers;
value.centers = newCenters;
LOG.info("old centers: " + oldCenters + ", new centers: " + newCenters);
// compare new/old centers
boolean converged = true;
for (int i = 0; i < value.centers.size() && converged; i++) {
Tuple oldCenter = (Tuple) oldCenters.get(i);
Tuple newCenter = (Tuple) newCenters.get(i);
double sum = 0.0d;
for (int j = 0; j < newCenter.size(); j++) {
double v = ((DoubleWritable) newCenter.get(j)).get()
- ((DoubleWritable) oldCenter.get(j)).get();
sum += v * v;
}
double dist = Math.sqrt(sum);
LOG.info("old center: " + oldCenter + ", new center: " + newCenter
+ ", dist: " + dist);
// converge threshold for each center: 0.05
converged = dist < 0.05d;
}
if (converged || context.getSuperstep() == context.getMaxIteration() - 1) {
// converged or reach max iteration, output centers
for (int i = 0; i < value.centers.size(); i++) {
context.write(((Tuple) value.centers.get(i)).toArray());
}
// true means to terminate iteration
return true;
}
// false means to continue iteration
return false;
}
}
private static void printUsage() {
System.out.println("Usage: <in> <out> [Max iterations (default 30)]");
System.exit(-1);
}
public static void main(String[] args) throws IOException {
if (args.length < 2)
printUsage();
GraphJob job = new GraphJob();
job.setGraphLoaderClass(KmeansVertexReader.class);
job.setRuntimePartitioning(false);
job.setVertexClass(KmeansVertex.class);
job.setAggregatorClass(KmeansAggregator.class);
job.addInput(TableInfo.builder().tableName(args[0]).build());
job.addOutput(TableInfo.builder().tableName(args[1]).build());
// default max iteration is 30
job.setMaxIteration(30);
if (args.length >= 3)
job.setMaxIteration(Integer.parseInt(args[2]));
long start = System.currentTimeMillis();
job.run();
System.out.println("Job Finished in "
+ (System.currentTimeMillis() - start) / 1000.0 + " seconds");
}
}
Configuration du job
La méthode main configure le GraphJob avec les paramètres suivants :
| Paramètre | Valeur | Description |
|---|---|---|
setGraphLoaderClass |
KmeansVertexReader.class |
Charge les enregistrements de la table d'entrée en tant que sommets |
setVertexClass |
KmeansVertex.class |
Définit la logique de calcul par sommet |
setAggregatorClass |
KmeansAggregator.class |
Définit la logique de mise à jour des centroïdes et de convergence |
setRuntimePartitioning |
false |
Désactive le partitionnement de graphe à l'exécution. Les sommets K-means n'ayant pas besoin d'être redistribués lors du chargement, la désactivation de cette option améliore les performances de chargement du graphe. |
setMaxIteration |
30 (par défaut) |
Définit la limite d'itérations. Transmettez un troisième argument au job pour remplacer cette valeur. |
addInput / addOutput |
args[0] / args[1] |
Noms des tables d'entrée et de sortie, transmis en tant qu'arguments de ligne de commande |