In [1]:
import numpy as np
In [2]:
A = {
    '+': {
        '+': 0.8,
        '-': 0.2
    },
    
    '-':{
        '+': 0.1,
        '-': 0.9
    }
}

B = {
    '+': {
        'A': 0.2,
        'T': 0.2,
        'C': 0.3,
        'G': 0.3
    },
    
    '-': {
        'A': 0.4,
        'T': 0.4,
        'C': 0.1,
        'G': 0.1
    }
}

P = {
    '+': 0.5,
    '-': 0.5
}
In [3]:
# HMM sa diskretnim raspodelama emisija
class HMM:
    def __init__(self, l = None):        
        if l != None:
            A, B, P = l
            self.A = A
            self.B = B
            self.P = P
        
    def a(self, q_1, q):
        return self.A[q_1][q]
    
    def b(self, q, x):
        return self.B[q][x]
    
    def pi(self, q):
        return self.P[q]
    
    def state_num(self, q):
        return list(self.A.keys()).index(q)
    
    def num_state(self, num):
        return list(self.A.keys())[num]
    
    def viterbi(self, X):
        T = len(X)
        N = len(self.A)

        v_matrix = [[0 for _ in range(T)] for _ in range(N)]
        backtrack_matrix = [[-1 for _ in range(T)] for _ in range(N)]
    
        for t in range(T):
            x = X[t]
            
            if t == 0:
                for i in range(N):
                    q = self.num_state(i)
                    
                    transition_prob = self.pi(q)
                    emission_prob = self.b(q, x)
                    
                    prob = transition_prob * emission_prob
                    
                    v_matrix[i][t] = prob
                    
            else:
                for i in range(N):
                    
                    max_prob = 0
                    max_prop_state = -1
                    
                    q = self.num_state(i)
                    
                    emission_prob = self.b(q, x)
                    
                    for j in range(N):
                        q_1 = self.num_state(j)
                        
                        prev_prob = v_matrix[j][t - 1]
                        
                        transition_prob = self.a(q_1, q)
                        
                        prob = transition_prob * emission_prob * prev_prob
                        
                        if prob > max_prob:
                            max_prob = prob
                            max_prop_state = j
                            
                    v_matrix[i][t] = max_prob
                    backtrack_matrix[i][t] = max_prop_state
                    
        
        # Rekonstrukcija puta
        last_index = np.argmax(np.array(v_matrix)[:,t - 1])
        
        path = ''
        t = T - 1
        
        while last_index != -1:
            last_state = self.num_state(last_index)
            path += last_state
            
            last_index = backtrack_matrix[last_index][t]
            t -= 1
            
        return ''.join(list(reversed(path)))
    
    def forward(self, X, k = None):
        if k == None:
            T = len(X)
        else:
            T = k
        
        N = len(self.A)

        v_matrix = [[0 for _ in range(T)] for _ in range(N)]
        
        for t in range(T):
            x = X[t]
            
            if t == 0:
                for i in range(N):
                    q = self.num_state(i)
                    
                    transition_prob = self.pi(q)
                    emission_prob = self.b(q, x)
                    
                    prob = transition_prob * emission_prob
                    
                    v_matrix[i][t] = prob
                    
            else:
                for i in range(N):

                    sum_prob = 0

                    q = self.num_state(i)

                    emission_prob = self.b(q, x)

                    for j in range(N):
                        q_1 = self.num_state(j)

                        prev_prob = v_matrix[j][t - 1]

                        transition_prob = self.a(q_1, q)

                        prob = transition_prob * emission_prob * prev_prob

                        sum_prob += prob

                    v_matrix[i][t] = sum_prob
                
        m = np.array(v_matrix)
        
        return m[:,T - 1].sum(), m[:,T - 1], m
    
    def backward(self, X, k = None):
        T = len(X)
        N = len(self.A)
        
        if k == None:
            start = 0
        else:
            start = k

        v_matrix = [[0 for _ in range(T)] for _ in range(N)]
        
        for t in reversed(range(start, T)):
            x = X[t]
            
            if t == T - 1:
                for i in range(N):
                    q = self.num_state(i)
                    
                    transition_prob = 1
                    emission_prob = self.b(q, x)
                    
                    prob = transition_prob * emission_prob
                    
                    v_matrix[i][t] = prob
                    
            else:
                for i in range(N):

                    sum_prob = 0

                    q = self.num_state(i)

                    emission_prob = self.b(q, x)

                    for j in range(N):
                        q_1 = self.num_state(j)

                        prev_prob = v_matrix[j][t + 1]

                        transition_prob = self.a(q, q_1)

                        prob = transition_prob * emission_prob * prev_prob
                        
                        if t == k:
                            prob *= self.pi(q)

                        sum_prob += prob

                    v_matrix[i][t] = sum_prob
                
        m = np.array(v_matrix)
        
        return m[:,k].sum(), m[:,k], m
    
    def baum_welch_single_sequence(self, X):
        _, _, all_alpha = self.forward(X)
        _, _, all_beta = self.backward(X)
        
        T = len(X)
        N = len(self.A)
        
        gamma = all_alpha * all_beta
        
        for t in range(T):
            marg_prob = gamma[:, t].sum()
            gamma[:, t] /= marg_prob
            
        zeye = np.array([[[0.0 for t in range(T)] for j in range(N)] for i in range(N)])
        
        for t in range(T - 1):
            marg_prob = 0
            for i in range(N):
                qi = self.num_state(i)
                
                for j in range(N):
                    qj = self.num_state(j)
                    
                    prob = all_alpha[i, t] * self.a(qi, qj) * all_beta[j, t + 1] * self.b(qj, X[t + 1])
                    zeye[i, j, t] = prob
                    marg_prob += prob
                    
            zeye[:,:,t] /= marg_prob
            
        new_P = {}
        
        for i in range(N):
            qi = self.num_state(i)
            
            new_P[qi] = gamma[i, 0]
            
        new_a = {}
        
        for i in range(N):
            qi = self.num_state(i)
            
            if qi not in new_a:
                new_a[qi] = {}
            
            for j in range(N):
                qj = self.num_state(j)
                
                new_a[qi][qj] = zeye[i,j, : T - 1].sum() / gamma[i, : T - 1].sum()
                
        v = ['A','T', 'C', 'G']
        
        new_b = {}
        
        for i in range(N):
            qi = self.num_state(i)
            new_b[qi] = {}
            
            for vk in v:
                indicator = np.array([int(X[t] == vk) for t in range(T)])
                
                new_b[qi][vk] = (gamma[i,:] * indicator).sum() / gamma[i,:].sum()
            
            
        return new_a, new_b, new_P
    
    def x_prob(self, X_arr):
        total_prob = 1.0
        
        for x in X_arr:
            total_prob *= self.forward(x)[0]
            
        return total_prob
    
    def baum_welch(self, X_arr):
        R = len(X_arr)
        N = len(self.A)
        
        eps = pow(10, -10)
            
        old_prob = 0
        new_prob = 1
        
        v = ['A','T','C','G']
        
        while True:
            old_prob = self.x_prob(X)
            
            a = []
            b = []
            p = []
            
            for x in X_arr:
                ai, bi, pi = self.baum_welch_single_sequence(x)
                a.append(ai)
                b.append(bi)
                p.append(pi)
                
            new_P = {}
            new_A = {}
            new_B = {}
                
            for r in range(R):
                # P
                for i in range(N):
                    qi = self.num_state(i)
                    
                    if qi not in new_P:
                        new_P[qi] = 0
                        
                    new_P[qi] += (p[r][qi] / R)
                    
                # A
                for i in range(N):
                    qi = self.num_state(i)
                    for j in range(N):
                        qj = self.num_state(j)

                        if qi not in new_A:
                            new_A[qi] = {}
                            
                        if qj not in new_A[qi]:
                            new_A[qi][qj] = 0

                        new_A[qi][qj] += (a[r][qi][qj] / R)
                        
                # B
                for i in range(N):
                    qi = self.num_state(i)
                    for vk in v:

                        if qi not in new_B:
                            new_B[qi] = {}
                            
                        if vk not in new_B[qi]:
                            new_B[qi][vk] = 0

                        new_B[qi][vk] += (b[r][qi][vk] / R)
                

            self.A = new_A
            self.P = new_P
            self.B = new_B
                
            new_prob = self.x_prob(X)
            
            if new_prob <= old_prob + eps:
                break
        
    
    def forward_backward(self, X, t):
        _, alpha, _ = self.forward(X, t)
        _, beta, _ = self.backward(X, t)
        prod = (alpha * beta)
        norm_prod = prod / prod.sum()
        
        return norm_prod
In [4]:
l = (A, B, P)
hmm = HMM(l)

X = 'GGCCTGATTATATTA'

# res = hmm.viterbi(X)
# res_f = hmm.forward(X, 8)
# res_b = hmm.backward(X, 8)
# hmm.baum_welch([X])

# t = 2
# res_fb = hmm.forward_backward(X, t)

# print(X)
# print(res_fb)