Volver a la lista

Baum-Welch EM: algoritmo iterativo que aprende automáticamente los parámetros de un HMM únicamente a partir de las observaciones

¿Cómo se aprende la tabla de probabilidades de un HMM cuando solo hay secuencias observadas y no hay etiquetas de estado? Desde la derivación de las dos etapas del algoritmo EM hasta el bucle de convergencia en Python y las trampas de los óptimos locales.

Avanzado
|
22min
|
Verificado (2026-07-20)
hidden markov modelunsupervised learningparameter estimation
Progreso0/120 (0%)

¿De dónde provienen las tablas de probabilidad?

En M18 y M19 comenzamos proporcionando aleatoriamente una tabla de probabilidad (π,a,b)(\pi, a, b). En la realidad, esta tabla no existe desde el principio. La situación más común es que solo haya secuencias observadas y no haya etiquetas de estado. ¿Cómo podemos aprender esta tabla a partir de los datos?

Si tuviéramos aprendizaje supervisado con respuestas correctas, bastaría con un simple conteo: contar cuántas veces aparece cada carácter en cada estado. Pero, ¿qué ocurre si no podemos observar los estados? El algoritmo Baum-Welch aborda este problema. Como caso especial del algoritmo Expectation-Maximization (EM), estima probabilísticamente la secuencia de estados mientras reestima iterativamente los parámetros. Este capítulo es el corazón del autoaprendizaje de la serie sobre HMM.

Las dos etapas de EM — Ajuste intuitivo

El proceso de EM funciona así:

  1. Inicialización: Asignar cualquier valor a la tabla de parámetros (completamente aleatorio o uniforme).
  2. Paso E (Expectation): Calcular las probabilidades γ\gamma para cada tiempo y estado, y las probabilidades de transición ξ\xi, utilizando los parámetros actuales.
  3. Paso M (Maximization): Tratar esas probabilidades como si fueran conteos reales para reestimar los parámetros.
  4. Repetir E y M hasta que la verosimilitud logarítmica converja.

La idea clave es esta: como no conocemos con certeza el estado, lo tratamos mediante probabilidades; y usamos esas probabilidades como si fueran conteos reales para recalcular los parámetros. Es un proceso de autoarranque (bootstrapping) que ajusta ligeramente el modelo hacia una mejor precisión.

Paso E — Extraer dos valores usando Forward-Backward

Los dos valores calculados en M19 sirven aquí como insumos.

γt(s)\gamma_t(s) = probabilidad posterior de estar en el estado ss en el tiempo tt

\gamma_t(s) = \frac{\alpha_t(s) \beta_t(s)}{P(O)}

ξt(i,j)\xi_t(i, j) = probabilidad posterior de transitar del estado ii en el tiempo tt al estado jj en el tiempo t+1t+1 (nuevo)

\xi_t(i, j) = \frac{\alpha_t(i) \cdot a(i \to j) \cdot b(j, o_{t+1}) \cdot \beta_{t+1}(j)}{P(O)}

ξ\xi es la versión "de transición" de γ\gamma. Combina las probabilidades forward (α\alpha) y backward (β\beta) de dos tiempos consecutivos para determinar la probabilidad de pasar de un estado a otro. Con estos dos valores, el paso M se convierte en un proceso similar al conteo.

Paso M — Reestimación por conteo

Reestimación de las probabilidades iniciales:

\pi'(s) = \gamma_1(s)

La probabilidad posterior en el primer tiempo es la nueva probabilidad inicial.

Reestimación de las probabilidades de transición:

a'(i \to j) = \frac{\sum_{t=1}^{T-1} \xi_t(i, j)}{\sum_{t=1}^{T-1} \gamma_t(i)}

El numerador es el "número esperado de transiciones de i a j", y el denominador es el "número esperado total de transiciones desde i hacia cualquier estado". La razón entre estos dos conteos corresponde directamente a la definición de probabilidad condicional.

Reestimación de las probabilidades de emisión:

b'(s, v_k) = \frac{\sum_{t: o_t = v_k} \gamma_t(s)}{\sum_{t=1}^{T} \gamma_t(s)}

El numerador es el "número esperado de observaciones del carácter vkv_k cuando se encuentra en el estado ss", y el denominador es el "número total esperado de veces que se visita el estado ss". La estructura es idéntica a la de los conteos.

Estas tres fórmulas constituyen todo el paso M (M-step). Simplemente se reestiman las estimaciones de máxima verosimilitud de los parámetros tratando γ\gamma y ξ\xi, obtenidos en el paso E (E-step), como si fueran conteos directos.

¿Por qué converge esto?

El algoritmo EM tiene una garantía sólida: la verosimilitud logarítmica de los datos aumenta monótonamente en cada iteración. Nunca empeora. Convergirá después de un número finito de iteraciones.

La demostración requiere la desigualdad de Jensen y la derivación de la función Q, lo cual excede el alcance de esta sección. Sin embargo, se puede intuir su funcionamiento. El paso E utiliza la distribución posterior generada por los parámetros actuales, y el paso M extrae los parámetros óptimos bajo dicha posterior. Cada etapa es esencialmente una versión probabilística del descenso de coordenadas, donde "se optimiza solo un lado mientras el otro permanece fijo".

Sin embargo, hay un punto de precaución: EM converge a un óptimo local, no garantiza el óptimo global. Si los valores iniciales son inadecuados, puede quedar atrapado en un óptimo local deficiente. En la práctica, se ejecuta múltiples veces con diferentes valores iniciales para seleccionar el mejor resultado, o se utilizan técnicas de agrupamiento previo como K-means para fijar los valores iniciales.

Realicemos una iteración manualmente

Utilizamos directamente α\alpha, β\beta y γ\gamma calculados en M19. La observación es GCAT, con la tabla de probabilidades inicial idéntica a la de M18.

La posterior de M19 (γ\gamma):

tγ(H)\gamma(H)γ(L)\gamma(L)
10.8370.163
20.8160.184
30.2500.749
40.1990.801

(Los valores de γ\gamma para t=1,2t=1, 2 se pueden calcular manualmente mediante la fórmula α×β/P(O)\alpha \times \beta / P(O); esto queda como ejercicio de verificación.)

Ejemplo de cálculo de ξ\xi (t=1t=1, transición H→H): \xi_1(H, H) = \frac{\alpha_1(H) \cdot a(H\to H) \cdot b(H, C) \cdot \beta_2(H)}{P(O)} = \frac{0.20 \times 0.7 \times 0.4 \times 0.0469}{0.003676} = \frac{0.002627}{0.003676} = 0.715

De esta manera, se obtienen 12 valores de ξ\xi: 4 combinaciones (i,j)(i, j) multiplicadas por 3 instantes temporales. En la reestimación de las transiciones...

a'(H \to H) = \frac{\xi_1(H,H) + \xi_2(H,H) + \xi_3(H,H)}{\gamma_1(H) + \gamma_2(H) + \gamma_3(H)}

Al finalizar este cálculo se obtiene una nueva tabla de probabilidades de transición, que probablemente difiere ligeramente del valor original de 0.7. Se utiliza esta nueva tabla para ejecutar la siguiente etapa E y luego la etapa M. Este proceso se repite hasta que la verosimilitud logarítmica se estabilice dentro del margen de error por redondeo.

Implementación en Python (incluyendo bucle de convergencia)

python
import math
def logsumexp(vals):
m = max(vals)
return m + math.log(sum(math.exp(v - m) for v in vals)) if m != float('-inf') else float('-inf')
def forward_backward(obs, N, start_p, trans_p, emit_p):
T = len(obs)
log = math.log
alpha = [[log(start_p[s]) + log(emit_p[s][obs[0]]) if t == 0 else 0.0
for s in range(N)] for t in range(T)]
for t in range(1, T):
for s in range(N):
terms = [alpha[t-1][sp] + log(trans_p[sp][s]) for sp in range(N)]
alpha[t][s] = logsumexp(terms) + log(emit_p[s][obs[t]])
log_P_O = logsumexp(alpha[T-1])
beta = [[0.0] * N for _ in range(T)]
for t in range(T-2, -1, -1):
for s in range(N):
terms = [log(trans_p[s][sp]) + log(emit_p[sp][obs[t+1]]) + beta[t+1][sp]
for sp in range(N)]
beta[t][s] = logsumexp(terms)
return alpha, beta, log_P_O
def baum_welch(obs, N, M, start_p, trans_p, emit_p, max_iter=100, tol=1e-5):
T = len(obs)
prev_ll = float('-inf')
for iteration in range(max_iter):
alpha, beta, log_P_O = forward_backward(obs, N, start_p, trans_p, emit_p)
# E-step
gamma = [[math.exp(alpha[t][s] + beta[t][s] - log_P_O) for s in range(N)]
for t in range(T)]
xi = [[[0.0] * N for _ in range(N)] for _ in range(T-1)]
for t in range(T-1):
for i in range(N):
for j in range(N):
num = alpha[t][i] + math.log(trans_p[i][j]) + \
math.log(emit_p[j][obs[t+1]]) + beta[t+1][j]
xi[t][i][j] = math.exp(num - log_P_O)
# M-step
start_p = [gamma[0][s] for s in range(N)]
for i in range(N):
denom = sum(gamma[t][i] for t in range(T-1))
for j in range(N):
trans_p[i][j] = sum(xi[t][i][j] for t in range(T-1)) / denom
for s in range(N):
denom = sum(gamma[t][s] for t in range(T))
for k in range(M):
num = sum(gamma[t][s] for t in range(T) if obs[t] == k)
emit_p[s][k] = num / denom
# Comprobar la convergencia
if abs(log_P_O - prev_ll) < tol:
break
prev_ll = log_P_O
return start_p, trans_p, emit_p, log_P_O, iteration + 1
# Empezar con valores iniciales aleatorios
import random
random.seed(42)
N, M = 2, 4
start_p = [0.6, 0.4]
trans_p = [[0.5, 0.5], [0.5, 0.5]]
emit_p = [[random.random() for _ in range(M)] for _ in range(N)]
for s in range(N):
total = sum(emit_p[s])
emit_p[s] = [v/total for v in emit_p[s]]
# Observaciones (varias concatenadas)
obs = [{'A':0,'C':1,'G':2,'T':3}[c] for c in 'GCATGCATGGCCAATT']
sp, tp, ep, ll, iters = baum_welch(obs, N, M, start_p, trans_p, emit_p)
print(f'Iteraciones hasta converger: {iters}, log P(O) final: {ll:.4f}')
print('Transiciones aprendidas:', tp)
print('Emisiones aprendidas:', ep)

Al ejecutarlo, converge en unas decenas de iteraciones y produce una tabla de probabilidades de emisión próxima a estados ricos en GC y ricos en AT. Sin haber recibido etiquetas correctas, el algoritmo descubre dos estados latentes únicamente a partir de las estadísticas de los caracteres. Este es el prototipo del aprendizaje no supervisado.

La trampa del óptimo local: algo que debes tener en cuenta

Si se ejecuta Baum-Welch varias veces, los resultados cambian en cada ocasión.

  • Con valores iniciales inadecuados, los dos estados convergen a distribuciones de emisión casi idénticas (colapso a un solo estado).
  • Si los valores iniciales resultan opuestos a la estructura de los datos, el algoritmo queda atrapado en un mal óptimo local.
  • El algoritmo puede detenerse satisfecho aunque la log-verosimilitud sea baja.

En la práctica, el estándar consiste en repetir la ejecución desde 20~100 puntos iniciales aleatorios y elegir la log-verosimilitud más alta. Otra opción es agrupar previamente las observaciones con K-means para obtener una tabla inicial de emisiones plausible. Sin estas técnicas, Baum-Welch no resulta estable en aplicaciones reales.

Complejidad

Una iteración de EM cuesta O(T × N²), igual que Forward-Backward. La convergencia suele requerir decenas o centenares de iteraciones, por lo que el coste total es O(T × N² × I). Sigue siendo práctico con T de decenas de miles; cuando existen varias secuencias, la paralelización por lotes es natural.

Correspondencia en informática: la esencia de EM

Esta es la versión probabilística del descenso por coordenadas y el bootstrapping mencionados en DryBench.

  • Descenso por coordenadas: optimiza alternativamente los parámetros y las variables latentes.
  • Bootstrapping: reutiliza como recuentos las probabilidades posteriores generadas por el propio modelo.
  • Convergencia monótona: la log-verosimilitud nunca disminuye.
  • Óptimo local: depende de los valores iniciales y requiere varios reinicios.

La idea de este algoritmo reaparece en otros contextos: EM para GMM (Gaussian Mixture Model), EM variacional para LDA (modelos de temas) y pseudoetiquetado en aprendizaje semisupervisado. Estimar la parte latente a partir de lo observado y reutilizar esa estimación como recuentos para volver a aprender los parámetros es un patrón que vertebra el aprendizaje de modelos probabilísticos.

Obtener una puntuación en Rosalind

El problema BW de Rosalind proporciona una secuencia observada y parámetros iniciales, y pide los parámetros finales tras N iteraciones de Baum-Welch. Basta con sustituir el número de iteraciones del bucle for anterior por el valor indicado en el ejercicio.

Ramificaciones para el siguiente artículo

  • Próximo artículo (M21): Profile HMM + detección de genes — aplica los HMM a la predicción génica real. En lugar de dos estados, utiliza decenas, como exones, intrones y sitios de empalme. Herramienta: HMMER.
  • Extensión relacionada: EM en línea — Actualización de parámetros en streaming sin tener que recorrer todos los datos cada vez. Esencial para el aprendizaje con secuencias de gran tamaño.
  • Aprendizaje de CRF — Cálculo directo del gradiente de la verosimilitud logarítmica para el aprendizaje de parámetros. Más flexible que HMM, pero computacionalmente más costoso.

Si desea profundizar más

  • Baum, Petrie, Soules, Weiss (1970), A maximization technique occurring in the statistical analysis of probabilistic functions of Markov chains, Annals of Mathematical Statistics — El artículo original. Aunque han pasado medio siglo, su integridad permanece intacta.
  • Dempster, Laird, Rubin (1977), Maximum Likelihood from Incomplete Data via the EM Algorithm, JRSS-B — El artículo que estableció el marco general del algoritmo EM.
  • UC Berkeley CS176 — Clase de aprendizaje de HMM del profesor Yun S. Song. Se observa la derivación de la función Q mediante anotaciones en pizarra.
  • Bishop Pattern Recognition and Machine Learning Capítulo 9 — Deriva GMM y EM de HMM integrados en un único marco.
  • Durbin et al. Biological Sequence Analysis Capítulo 3.3 — Tutorial de Baum-Welch en contexto biológico.

Antes de pasar al siguiente capítulo, ejecute el código de Python anterior y cambie los valores iniciales varias veces. La verdadera práctica de este capítulo consiste en adquirir intuición sobre qué valores iniciales convergen a un óptimo local bueno y cuáles quedan atrapados en uno malo.

💬 Preguntas y comentarios

0 comentarios

Puedes publicar sin iniciar sesión. Los comentarios de invitados no pueden editarse ni eliminarse después.

0/2000

Cargando...