Tous les produits
Search
Centre de documentation

MaxCompute:K-means clustering

Dernière mise à jour :Aug 10, 2026

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

  1. Sélectionnez les centroïdes initiaux pour les k clusters.

  2. 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.

  3. Recalculez chaque centroïde comme la moyenne arithmétique de tous les points attribués à ce cluster.

  4. 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ïdes

  • sums : sommes accumulées des coordonnées par cluster

  • counts : 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