Todos os produtos
Search
Central de documentação

MaxCompute:K-means clustering

Última atualização: Jun 26, 2026

A clusterização K-means divide um conjunto de dados em k grupos e refina iterativamente os centroides dos clusters até que se estabilizem ou atinjam o limite de iterações. Este tópico demonstra como implementar a clusterização K-means com a API Graph do MaxCompute (Java).

Funcionamento

Cada iteração atribui todos os pontos de dados ao centroide mais próximo e recalcula o centroide como a média de todos os pontos daquele cluster. O algoritmo para quando todos os centroides se deslocam menos que o limiar de convergência ou quando atingem o número máximo de iterações.

Etapas do algoritmo

  1. Selecione os centroides iniciais para k clusters.

  2. Calcule a distância euclidiana quadrada de cada ponto de dados em relação a todos os k centroides e atribua-o ao cluster mais próximo.

  3. Recalcule cada centroide como a média aritmética de todos os pontos atribuídos àquele cluster.

  4. Se todos os centroides se deslocarem menos que o limiar de convergência, pare. Caso contrário, repita a partir da etapa 2.

Parâmetros principais

Parâmetro

Padrão

Descrição

Limiar de convergência

0,05

Distância euclidiana máxima que um centroide pode se deslocar entre iterações antes que o algoritmo considere a convergência. O cálculo usa a raiz quadrada da soma das diferenças de coordenadas ao quadrado.

Máximo de iterações

30

Limite superior para o número de supersteps. O job é interrompido mesmo que os centroides não tenham convergido totalmente.

Métrica de distância

Euclidiana quadrada

A distância quadrada serve para comparar atribuições de clusters. A raiz quadrada é calculada apenas durante a verificação de convergência em terminate().

Implementação

O exemplo a seguir implementa a clusterização K-means com quatro classes baseadas na API Graph do MaxCompute.

Visão geral dos componentes

Antes de analisar o código, entenda a responsabilidade de cada componente:

**KmeansVertex** representa um único ponto de dados como uma Tupla (vetor de características). Seu método compute() passa o valor do vértice para o agregador ao chamar context.aggregate().

**KmeansVertexReader** carrega a tabela de entrada e cria um vértice por registro. O número do registro torna-se o ID do vértice; todas as colunas do registro formam o valor do vértice como uma Tupla.

**KmeansAggrValue** mantém o estado de agregação compartilhado entre os workers e contém três campos do tipo Tuple:

  • centers: coordenadas atuais dos centroides

  • sums: somas acumuladas de coordenadas por cluster

  • counts: contagem de pontos de dados por cluster

**KmeansAggregator** encapsula a lógica principal do algoritmo em quatro métodos:

Método

Responsabilidade

createInitialValue()

No superstep 0, lê os centroides iniciais do arquivo de cache chamado "centers". Nos supersteps seguintes, recupera os centroides calculados na iteração anterior por meio de getLastAggregatedValue(0).

aggregate()

Para cada vértice, encontra o centroide mais próximo usando a distância euclidiana quadrada e atualize os valores de sum e count daquele cluster.

merge()

Combina os valores parciais de sums e counts coletados dos workers em execução paralela.

terminate()

Calcula novos centroides a partir de sums e counts. Se todos os centroides se deslocarem menos que 0,05 ou o número máximo de iterações for atingido, grava os centroides finais na tabela de saída e retorna true (parar). Caso contrário, retorna false (continuar).

Código de exemplo

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");
  }
}

Configuração do job

O método main configura o GraphJob com as seguintes definições:

Configuração

Valor

Descrição

setGraphLoaderClass

KmeansVertexReader.class

Carrega os registros da tabela de entrada como vértices

setVertexClass

KmeansVertex.class

Define a lógica de computação por vértice

setAggregatorClass

KmeansAggregator.class

Define a lógica de atualização de centroides e convergência

setRuntimePartitioning

false

Desativa o particionamento do grafo em tempo de execução. Os vértices do K-means não precisam de redistribuição durante o carregamento; desativar essa opção melhora o desempenho do carregamento do grafo.

setMaxIteration

30 (padrão)

Defina o limite de iterações. Passe um terceiro argumento para o job para substituir esse valor.

addInput / addOutput

args[0] / args[1]

Nomes das tabelas de entrada e saída, passados como argumentos de linha de comando