function [newa, newpi, newb] = hmmlearn(a, pi, b, o, alpha, beta)

spi = size(pi');
so = size(o');
sb = size(b);

T = so(1);
N = spi(1);
nst = sb(2);

disp('Calculating zetas');
for t = 1:(T-1)
	bottom = 0;
	for i = 1:N
		for j = 1:N
			bottom = bottom + alpha(t, i) * a(i,j) * b(j,o(t+1)) * beta(t+1, j);
		end
	end
	for i = 1:N
		for j = 1:N
			if(bottom == 0)
				zeta(t,i,j) = 0;
			else
				zeta(t, i, j) = (alpha(t,i) * a(i,j) * b(j, o(t+1)) * beta(t+1, j))/bottom;
			end
		end
	end
end
disp('Calculating gammas');
for t = 1:(T-1)
	for i = 1:N
		gamma(t, i) = 0;
		for j = 1:N
			gamma(t,i) = gamma(t,i) + zeta(t, i, j);
		end
	end
end

disp('Updating a and pi');
for j = 1:N
	newpi(j) = gamma(1, j);
	for i = 1:N
		topsum = 0;
		botsum = 0;
		for t = 1:(T-1)
			topsum = topsum + zeta(t, i, j);
			botsum = botsum + gamma(t,i);
		end
		if(botsum == 0)
			newa(i,j) = 0;
		else
			newa(i,j) = topsum/botsum;
		end
	end
end

disp('Updating b');
for j = 1:N
	for k = 1:nst
		botsum = 0;
		topsum = 0;
		for t = 1:(T-1)
			botsum = botsum + gamma(t,j);
			if(o(t) == k)
				topsum = topsum + gamma(t,j);
			end
		end
		if(botsum == 0)
			newb(j,k) = 0;
		else
			newb(j,k) = topsum/botsum;
		end
	end
end

