bp_net.cpp

来自「故障诊断工作涉及的领域相当广泛」· C++ 代码 · 共 1,535 行 · 第 1/3 页

CPP
1,535
字号
									
								}while(count<lth);
failureinfo:
								while(c!='{'&&count<lth)
								ar>>c;
								do{
									do{
									ar>>c;
									*(bag++)=c;	
									count++;//字符计数
									}while(c!='}'&&c!='\n'&&c!=EOF);
									if(bag0<bag-2)
									{failuretype[index].Empty();
								for(unsigned i=0;i<unsigned(bag-bag0);i++)
									failuretype[index]+=bag0[i];
								index++;
								bag=bag0;
										memset(bag,0,sizeof(char)*64);
								if(index==nout)
									goto tail;
									}
								}while(count<lth);
										
							tail:	
								ar.Close();
								filebin->Close();
								;
							}
					}
				
					unsigned nsamp1=incount/nin;
					unsigned nsamp2=outcount/nout;
					if(nsamp1==nsamp2&&nsamp1>0)
						{	nsamp=nsamp1;
							realloc(*ibag,(incount+outcount)*sizeof(double));
							nscnt=incount;
							nsdocnt=outcount;
							if(nsamp>0)
								pvin=*ibag;
							mwin=mwArray(nin,nsamp,pvin);
							pvout=new double[nout*nsamp];
							pvstdout=*ibag+nin*nsamp;
							mwstdout=mwArray(nout,nsamp,pvstdout);
							bvalid=TRUE;
							tterror=new double[nsamp];
							merror=new mwArray[nsamp];
							
	//	for(i=0;i<nsamp;i++)
		//	merror[i]=mwArray(1);
		
	//	foward();			
							if(index==nout)
							{bvalid=TRUE;
							return TRUE;
							}
							else
								return FALSE;
							}
							else 
							{
								static char info[]="样本数据数量与标准输出信息数量不匹配";
								*errorinfo=info;
								return FALSE;
							}
				}
static char info[]="样本数据数量与标准输出信息数量不匹配";
*errorinfo=NULL;
return FALSE;
}

BOOL bp_net::diagnose(const CString** result,unsigned int *cnt,unsigned int* pidx,double* perror)
{	*cnt=0;
	int l,c,index;double value=100,dbag;
	mwArray mbag,error,caltbar,caltbar1;
	char* errorinfo;
	*result=failuretype;
	memset(bag,0,sizeof(double)*100);
	if(read_data(&pdata,&errorinfo))//获得数据
	{			
	mdout=mdata;
	for(unsigned int i=0;i<nlayer;i++)
		{	mdout.ExtractData(bag);
			memset(bag,0,sizeof(double)*100);
			player[i].indata=mdout;
			mwArray oneone=ones(mwVarargin(player[i].mbar.Size(1)),mdata.Size(2));
			for(unsigned int j=0;j<mdata.Size(2);j++)
				caltbar=horzcat(caltbar,player[i].mbar);
			mdout=power(exp(-player[i].weight*mdout+caltbar)+oneone,-1);
			player[i].outdata=mdout;
			mdout.ExtractData(bag);	
			l=mdout.Size(1);
			c=mdout.Size(2);
			memset(bag,0,sizeof(double)*100);
			caltbar=mwArray();
		}
		lastin=mdout;
			for(unsigned int j=0;j<mdata.Size(2);j++)
				caltbar1=horzcat(caltbar1,lastbar);
		mwArray oneone=ones(lastbar.Size(1),mdout.Size(2));
		switch(func_type)
		{case simulate:
			mdout=weightout*mdout+caltbar1;//现行输出层代码
			break;
		default:
			mdout=power(exp(-weightout*mdout+caltbar1)+oneone,-1);//sigmod型输出层代码
		}
		c=mdout.Size(2);
		mdout.ExtractData(bag);
		caltbar1=mwArray();
		for(j=0;j<c;j++)
		{memset(bag,0,sizeof(double)*100);
			for(i=0;i<nsamp;i++)
			{	mbag=mdout(colon(),int(j+1))-mwstdout(colon(),int(i+1));
				error=sqrt(sum(times(mbag,mbag)/(mwArray)(int)nout));//阵列乘法 
				dbag=error.ExtractScalar(1);
				if(value>dbag)
				{value=dbag;
				index=i;
				}
			
			}
		perror[*cnt]=value;
		pidx[(*cnt)++]=index;
		value=100;
		}
		return TRUE;
	}
	return FALSE;
}

void bp_net::foward(unsigned int ind)
{int	index=ind+1;
	int l,c;
	double bag[200];
	memset(bag,0,sizeof(double)*200);
	mwout=mwin(colon(),mwArray(int(index)));
	mwout.ExtractData(bag);	
	memset(bag,0,sizeof(double)*200);
	for(unsigned int i=0;i<nlayer;i++)
	{	
		player[i].weight.ExtractData(bag);
		memset(bag,0,sizeof(double)*200);
		player[i].mbar.ExtractData(bag);
		memset(bag,0,sizeof(double)*200);
		player[i].indata=mwout;
	//	mwArray oneone=ones(mwVarargin(player[i].mbar.Size(1)),1);	
	//	mwout=power(exp(-player[i].weight*mwout+player[i].mbar)+oneone,-1);
		player[i].calt(mwout);
		player[i].outdata=mwout;
		mwout.ExtractData(bag);	
		l=mwout.Size(1);
		c=mwout.Size(2);
		memset(bag,0,sizeof(double)*200);
	}
		lastin=mwout;
		mwout=weightout*mwout+lastbar;
	;
		mwout.ExtractData(bag);
		memset(bag,0,sizeof(double)*100);
		
}

void bp_net::backward(int ind)
{	int l,c;
	memset(bag,0,sizeof(double)*100);
	int index=ind+1;
	l=weightout.Size(1);
	mwArray oneone=ones(l,1);	
	mwArray mbag=weightout;
	mwArray mbag1=mwstdout(colon(),mwArray(int(index)));//当前样本的标准输出
	dedlastin=mwstdout(colon(),mwArray(int(index)))-mwout;
	//dedlastin=times(times(mwstdout(colon(),mwArray(int(index)))-mwout,mwout),oneone-mwout);
	mwArray howfar=dedlastin*rot90(lastin);
	howfar.ExtractData(bag);
    memset(bag,0,sizeof(double)*100);
	l=howfar.Size(1);
	c=howfar.Size(2);
	weightout=weightout+alpha*howfar+beta*(weightout-preweightout);
	preweightout=mbag;
	mbag=lastbar;
	lastbar=lastbar+alpha*dedlastin+beta*(lastbar-prelastbar);
//	lastbar=lastbar/sqrt(sum(times(lastbar,lastbar)));
	prelastbar=mbag;
	weightout.ExtractData(bag);
    memset(bag,0,sizeof(double)*100);
	player[nlayer-1].dedout=rot90(weightout)*dedlastin;
	c=dedlastin.Size(1);
	for(int i=nlayer-1;i>=0;i--)
	{
		l=player[i].outdata.Size(1);
		oneone=ones(l,1);
		mbag=times(times(player[i].dedout,player[i].outdata),oneone-player[i].outdata);
		howfar=mbag*rot90(player[i].indata);
		if(i>0)
			player[i-1].dedout=rot90(player[i].weight)*mbag;
		mbag=player[i].mbar;
		player[i].mbar=player[i].mbar-alpha*mbag+beta*(player[i].mbar-player[i].prembar);
		player[i].prembar=mbag;
		mbag=player[i].weight;
		player[i].weight=player[i].weight+alpha*howfar+beta*(player[i].weight-player[i].preweight);
		player[i].preweight=mbag;
		player[i].weight.ExtractData(bag);	
		memset(bag,0,sizeof(double)*100);
		
	}
int v=1;
}




BOOL bp_net::low_error()
{
return TRUE;
}
void bp_net::draw_area(CDC *pdc,BOOL visible)
{	CPen* pNewPen=new CPen;
	pNewPen->CreatePen(0,5,RGB(50,50,255));	
	CPen* oldpen=pdc->SelectObject(pNewPen);
	if(visible)
	{
		pdc->MoveTo(dipan.left,dipan.top);
		pdc->LineTo(dipan.left,dipan.bottom);
		pdc->LineTo(dipan.right,dipan.bottom);
		pdc->LineTo(dipan.right,dipan.top);
		pdc->LineTo(dipan.left,dipan.top);
	}
	else
	{	unsigned long bkcolor= pdc->GetBkColor();
		CPen* pNewPen=new CPen;
		pNewPen->CreatePen(0,5,bkcolor);	
		pdc->SelectObject(pNewPen);
		//覆盖先前的地盘
		pdc->MoveTo(dipan.left,dipan.top);
		pdc->LineTo(dipan.left,dipan.bottom);
		pdc->LineTo(dipan.right,dipan.bottom);
		pdc->LineTo(dipan.right,dipan.top);
		pdc->LineTo(dipan.left,dipan.top);
		frenew=FALSE;
	}
	pdc->SelectObject(oldpen);
}
void bp_net::renew_net(CDC *pdc,CPoint& point)
{
if(is_snap(point))
{
	draw_area(pdc,TRUE);
	bpnet_basic_info bpdlg(this);
	bpdlg.DoModal();
}
else if(!is_snap(point)&&frenew)
	draw_area(pdc,FALSE);
}

BOOL bp_net::is_snap(CPoint &point)
{
if(layin.is_snap(point))
	return FALSE;
for(unsigned int i=0;i<nlayer;i++)
{if(player[i].is_snap(point))
	return FALSE;
}
if(layout.is_snap(point))
	return FALSE;
if(dipan.PtInRect(point))
{	frenew=TRUE;
	return TRUE;
}
	return FALSE;
}


void bp_net::reset(double *databag)//重构权值和阈值矩阵
{double bag[200];
 memset(bag,0,sizeof(double)*200);
 unsigned where=0;
 int l=player->weight.Size(1);
 int c=player->weight.Size(2);
			player->weight=mwArray(l,c,databag);
			player->weight.ExtractData(bag);
			memset(bag,0,sizeof(double)*200);
			where=l*c;
			player->mbar=mwArray(l,1,databag+where);
			where+=l;
				for(unsigned int i=1;i<nlayer;i++)//重构隐含层及其权矩阵
				{l=(player+i)->weight.Size(1);c=(player+i)->weight.Size(2);
					(player+i)->weight=mwArray(l,c,databag+where);
					where+=l*c;
					(player+i)->mbar=mwArray(c,1,databag+where);
					where+=l;
				}
				//重构输出层及其权矩阵
			l=weightout.Size(1);c=weightout.Size(2);
			weightout=mwArray(l,c,databag+where);
			where+=l*c;
			lastbar=mwArray(l,1,databag+where);		
}

	//mwArray bag1=times(weightout,weightout);
	//	mwArray bag2,bag3;
	//	for(i=0;i<bag1.Size(1);i++)
	//	{	bag2=ones(1,bag1.Size(2));
	//		bag2=bag2*sum(bag1(int(i+1),colon()));
	//	if(i==0)
	//		bag3=bag2;
	//	else
	//		bag3=vertcat(bag3,bag2);
	//	}
	//	weightout=rdivide(weightout,bag3);
	//	weightout.ExtractData(bag);
	//	memset(bag,0,sizeof(double)*100);
	//	mwout.ExtractData(bag);
	//	memset(bag,0,sizeof(double)*100);
	//	lastbar.ExtractData(bag);
	////	memset(bag,0,sizeof(double)*100);
	//	mwout=mwout/sqrt(sum(times(mwout,mwout)))

BOOL bp_net::isvalid()
{
return bvalid;
}

BOOL bp_net::istrained()
{
return btrained;
}

void bp_net::init()
{
failuretype=NULL;
	type=std;
	mwArray m0=mwArray(0);
	bempty=TRUE;
	lastbar=m0;
	lastin=m0;
	mdata=m0;
	mdout=m0;
	mwin=m0;
	mwout=m0;
	mwstdout=m0;
	//mx=m0;
	prelastbar=m0;
	preweightout=m0;
	weightout=m0;
	bvalid=FALSE;
	nin=0;
	nout=0;
	nsamp=0;
	nlayer=0;
	player=NULL;
	ndcnt=0;
	nscnt=0;
	nsdocnt=0;
	pvin=NULL;
	pvout=NULL;
	pvstdout=NULL;
	tterror=NULL;//输出层总体误差初值
	merror=NULL;
	stderror=0.0001;//缺省输出层标准误差
//	outerror=NULL;//输出层节点误差
	alpha=0.22;////缺省学习因子
	beta=0.11;////缺省冲量因子
	cnt=0;
	bsucceed=FALSE;
	btrained=FALSE;
	pdataout=NULL;
	ibag=NULL;
	frenew=FALSE;
	//mx=mwArray(1);
}

const mwArray* bp_net::get_data()
{
return &mdata;
}

void bp_net::setmatrix(bp_net &in)
{

}

void bp_net::farward()
{
	int l,c;
	double bag[200];
	memset(bag,0,sizeof(double)*200);
	mwout=mwin;
	mwout.ExtractData(bag);	
	memset(bag,0,sizeof(double)*200);
	weightout.ExtractData(bag);	
	memset(bag,0,sizeof(double)*200);
	for(unsigned int i=0;i<nlayer;i++)
	{	
		player[i].weight.ExtractData(bag);
		memset(bag,0,sizeof(double)*200);
		player[i].mbar.ExtractData(bag);
		memset(bag,0,sizeof(double)*200);
		player[i].indata=mwout;
	//	mwArray oneone=ones(mwVarargin(player[i].mbar.Size(1)),1);	
	//	mwout=power(exp(-player[i].weight*mwout+player[i].mbar)+oneone,-1);
		player[i].calt(mwout);
		player[i].outdata=mwout;
		mwout.ExtractData(bag);	
		l=mwout.Size(1);
		c=mwout.Size(2);
		memset(bag,0,sizeof(double)*200);
	}
		lastin=mwout;
		mwArray mbag;
		for(i=0;i<mwout.Size(2);i++)
		mbag=horzcat(mbag,lastbar);
		mwArray oneone=ones(lastbar.Size(1),mwout.Size(2));
		switch(func_type)
		{case simulate:
			mwout=weightout*mwout+mbag;//现行输出层的代码
		 break;
		 default:
			mwout=power(exp(-weightout*mwout+mbag)+oneone,-1);//sigmod型输出层代码
		}
		mwout.ExtractData(bag);
		memset(bag,0,sizeof(double)*100);
}

void bp_net::backward()
{
	int l,c;
	memset(bag,0,sizeof(double)*100);
	l=weightout.Size(1);
	mwArray oneone=ones(l,1);	
	mwArray mbag=weightout,dwout,dbout,*dwp=new mwArray[nsamp],*dbp=new mwArray[nsamp];
	mwArray mbag1,mbag2;//当前样本的标准输出
	mbag1=mwstdout-mwout;
	for(int j=0;j<mbag1.Size(2);j++)
	{l=mwout.Size(1);
	oneone=ones(l,1);
	switch(func_type)
	{case simulate:
	dedlastin=mbag1(colon(),j+1);//现行输出层的代码
	break;
	default://sigmod型输出层的代码
	dedlastin=times(times(mwstdout(colon(),j+1)-mwout(colon(),j+1),mwout(colon(),j+1)),\
		oneone-mwout(colon(),j+1));
	}
	mwArray howfar=dedlastin*rot90(lastin(colon(),j+1));
	howfar.ExtractData(bag);
    memset(bag,0,sizeof(double)*100);
	l=howfar.Size(1);
	c=howfar.Size(2);
	if(j==0)
		dwout=alpha*howfar+beta*(weightout-preweightout);
	else
		dwout=dwout+alpha*howfar+beta*(weightout-preweightout);
	if(j==0)
		dbout=alpha*dedlastin+beta*(lastbar-prelastbar);
	else
		dbout=dbout+alpha*dedlastin+beta*(lastbar-prelastbar);
	player[nlayer-1].dedout=rot90(weightout)*dedlastin;
	c=dedlastin.Size(1);
	for(int i=nlayer-1;i>=0;i--)
	{
		l=player[i].outdata.Size(1);
		oneone=ones(l,1);
		mbag=times(times(player[i].dedout,player[i].outdata(colon(),j+1)),\
			oneone-player[i].outdata(colon(),j+1));
		howfar=mbag*rot90(player[i].indata(colon(),j+1));
		if(i>0)
			player[i-1].dedout=rot90(player[i].weight)*mbag;
		if(j==0)
			dbp[i]=-alpha*mbag+beta*(player[i].mbar-player[i].prembar);
		else
			dbp[i]=dbp[i]-alpha*mbag+beta*(player[i].mbar-player[i].prembar);
		if(j==0)
			dwp[i]=alpha*howfar+beta*(player[i].weight-player[i].preweight);
		else
			dwp[i]=dwp[i]+alpha*howfar+beta*(player[i].weight-player[i].preweight);		
	}
	}
	preweightout=weightout;
	weightout=weightout+dwout;
//	prelastbar=lastbar;
//	lastbar=lastbar+dbout;
	for(int i=nlayer-1;i>=0;i--)
	{	player[i].prembar=player[i].mbar;
		player[i].mbar=player[i].mbar+dbp[i];
		player[i].preweight=player[i].weight;
		player[i].weight=player[i].weight+dwp[i];
	}
	delete[] dwp;
	delete[] dbp;
int v=1;
}

double bp_net::error()
{	mwArray mbag1,mbag2;
int l,c;
l=mwout.Size(1);
c=mwout.Size(2);
l=mwstdout.Size(1);
c=mwstdout.Size(2);
	mbag1=mwstdout-mwout;
	mbag2.Clear();
	for(int u=0;u<mbag1.Size(2);u++)
		mbag2=vertcat(mbag2,mbag1(colon(),u+1));
	mwArray error=sum(times(mbag2,mbag2)/2);//阵列乘法 
	error.ExtractData(bag);	
	memset(bag,0,sizeof(double)*100);
	double value=error.ExtractScalar(1);
	return value;

}

⌨️ 快捷键说明

复制代码Ctrl + C
搜索代码Ctrl + F
全屏模式F11
增大字号Ctrl + =
减小字号Ctrl + -
显示快捷键?