/*
 * delta2.c
 *
 * compute deltas when matrix has been hand-configured
 */

#include <stdio.h>
#include "defs.h"

delta2(teachp) 
	float	*teachp;
{
	extern	float **d_tokenp;
	extern	int debug;
	extern	int numtotal;
	extern	int numout;
	extern	int numin;
	extern	int cur_token;
	extern	int rbp;
	extern	int bipolar;
	extern	int all_linear;
	extern	int pcomp;
	extern	float **t_tokenp;
	extern	float *nodep;
	extern	float *deltap;
	extern	float **weightp;
	extern	float global_error;		/* in delta1.c */
	extern	float *delta_z;
	extern	float eps;
	extern	float alpha;
	extern	float beta;
	extern	float length();
	extern	int **conup;
	register float *np;
	register float *dp;
	register float *tp;
	register float **wp;
	register float tmp;
	register int i;
	float 	I;
	float	J;
	float	diff;
	int	**cupp;
	int	*cup;
	int	bo;

	/*
	 * clear deltas
	 */
	lzero(deltap, numtotal);
	/* 
	 * compute deltas for output units
	 */
	if (rbp == 0) {
		global_error = 0;
		bo = numtotal - numout;
		for (i=bo, np=nodep+bo, dp=deltap+bo, tp=teachp; i<numtotal; i++,np++,dp++, tp++) {
			diff = *tp - *np;
			if (bipolar==0)
				*dp = diff * *np * (1 - *np);
			else
				*dp = .5 * diff * (1. + *np) * (1. - *np);
			if (all_linear){
				if ((*np > -10.) && (*np < 10.))
					*dp = diff;
				else
					*dp = 0.;
			}
			if (*tp <= -999999999.0)
				*dp = 0.0;
			global_error += diff*diff;
		}
	/*
	 * next compute deltas for all other units, working backwards.  If
	 * we are doing rbp, include the outputs in this as well, and
	 * iterate till achieve minimum delta_z.
	 */
		bo -= 1;
		wp = weightp;
		for (i=bo, np=nodep+bo, dp=deltap+bo, cupp=conup+bo; 
			i>0; i--, np--, dp--, cupp--) {
			tmp = 0.0;
			for (cup = *cupp; *cup != -1; cup++) {
				tmp += *(deltap+*cup) * *(*(wp+i)+*cup);
			}
			if (bipolar==0)
				*dp = (1. - *np) * *np * tmp;
			else
				*dp = .5 * (1. - *np) * (1. + *np) * tmp;
			if (all_linear){
				if ((*np > -10.) && (*np < 10.))
					*dp = tmp;
				else
					*dp = 0.;
			}
		}
	} else if (rbp == 1) {
		/*
		 * different than non-rbp case.  Logic is:
		 *   1. sum the prod. of (error of units we talk to *
		 *	deriv. of their act, * connection strenghts;
		 *   2. calc. delta_z: sum of (1) plus external error,
		 *	plus previous error
		 *   3. increment error by delta_z.
		 */
		if (debug > 5) {
		    fprintf(stdout, "A [%d]: ", cur_token);
		    for (i=0; i<numtotal; i++)
			fprintf(stdout, "%f ", *(nodep+i));
		    fprintf(stdout, "\n");
	        }
		do {
		    float act_r;
		    float der_r;
		    float g;
		    
		    if (debug > 5)
			fprintf(stdout, "J [%d]: ", cur_token);
		    wp = weightp;
		    global_error = 0.0;
		    for (i=0, dp=deltap, cupp=conup; 
			    i<numtotal; i++, dp++, cupp++) {
			    g = 0.0;
			    /*
			     * accumulate new error
			     */
			    for (cup = *cupp; *cup != -1; cup++) {
			        I = *(*(d_tokenp+cur_token) + *cup);
			        if (I <= -999999999.)
				    I = 0.0;
				act_r = (alpha * *(nodep + *cup) - I) / beta;
			        /*
				act_r = *(nodep + *cup);
				*/
				if (bipolar==0)
					der_r = (1-act_r) * act_r;
				else
					der_r = .5*(1.-act_r) * (1.+act_r);
				g += der_r *  *(*(wp+i)+*cup) * *(deltap+*cup);
			    }
			    /*
			     * if an output node, get target output and
			     * calculate target-actual difference
			     */
			    J = *(teachp+i);
			    if (J <= -999999999.0) {
				J = 0.0;
			    } else {
				J = J - *(nodep+i);
			    } 
			    if (debug > 5)
			    	fprintf(stdout, "%f ", J);
			    *(delta_z+i) = (-alpha * *dp) + (beta * g) + J;
			    *dp += *(delta_z+i);
		    }
		    if (debug > 5)
			fprintf(stdout, "\n");
		} while (length(delta_z,numtotal) > eps);
		global_error = 0.0;
		for (i=0; i<numtotal; i++) {
			J = *(teachp+i);
			if (J <= -99999999.0) {
				J = 0.0;
			} else {
				J = J - *(nodep+i);
				global_error += J*J;
			}
		}
		if (debug > 5) {
		    fprintf(stdout, "D [%d]: ", cur_token);
		    for (i=0; i<numtotal; i++)
			fprintf(stdout, "%f ", *(deltap+i));
		    fprintf(stdout, "\n");
	        }
	}
}

