Files
jens 82fbd5d19d Initial import
git-svn-id: http://moon:8086/svn/software/trunk/libsrc/nn@1 b431acfa-c32f-4a4a-93f1-934dc6c82436
2014-07-19 07:44:42 +00:00

114 lines
3.3 KiB
C
Executable File

/**********************************************************************
* nnet.h
*
* (C) 2000 J. Ahrensfeld
*
**********************************************************************/
#ifndef NNET_H
#define NNET_H
#define UTYPE_BIAS 0
#define UTYPE_INPUT 1
#define UTYPE_HIDDEN 2
#define UTYPE_OUTPUT 3
#define UNIT_ID_INPUT (-UTYPE_INPUT)
#define UNIT_ID_OUTPUT (-UTYPE_OUTPUT)
#ifndef M_PI
#define M_PI 3.1415926535897932384626433832795
#endif
#ifndef M_E
#define M_E 2.71828182845904523536028747135266
#endif
enum
{
NNTYPE_MLP = 0,
NNTYPE_RECURRENT
};
/**********************************************************************/
typedef struct _sWTINFO
{
UINT32 unitID, weightID;
} WTINFO;
typedef struct _sUNIT
{
UINT32 nWeights, nAxonTo, *pAxonFrom, UType, FType, globalUnitID, layerID, unitID;
WTINFO *pAxonTo;
FLOAT64 axon, lastAxon, *pWeight, *pdWeight, net, error, delta, *pSlope;
FLOAT64 (*F)(FLOAT64); /* Zeiger auf Aktivierungsfunktion */
FLOAT64 (*Fd)(FLOAT64); /* Zeiger auf Ableitungsfunktion */
FLOAT64 **ppRecurr, **ppRecurrOld;
} UNIT;
typedef struct _sLAYER
{
UINT32 nUnits, ID;
UNIT **ppUnit;
} LAYER;
typedef struct _sNNET
{
UINT32 nLayers, nUnits, nnType, ID, nIn, nOut, nHidden;
struct _sLAYER *pLayer;
struct _sUNIT *pUnit;
UINT32 *pInputID, *pHiddenID, *pOutputID;
FLOAT64 sse, initwrange;
} NNET;
enum
{
Null, Lin, Tanh, ASig, Sig, SSig, FSig
};
UINT32 NetInit(struct _sNNET *pObj, struct _sNETFILE *pNetData, UINT32 netID);
UINT32 NetAddUnit(struct _sNNET *pObj, UINT32 unitType, UINT32 FType);
UINT32 NetAddConnection(struct _sNNET *pObj, UINT32 ID_i, UINT32 ID_j);
UINT32 NetFeedForward(struct _sNNET *pObj, FLOAT64 *pIn, FLOAT64 *pOut);
UINT32 NetBackPropErr(struct _sNNET *pObj, FLOAT64 *pIn, FLOAT64 *pTarget);
UINT32 NetBackPropErr2(struct _sNNET *pObj, FLOAT64 *pIn, FLOAT64 *pError);
UINT32 NetWeightUpd(struct _sNNET *pObj, FLOAT64 *pIn, UINT32 len);
UINT32 NetTrain(struct _sNNET *pObj, FLOAT64 *pIn, FLOAT64 *pOut, FLOAT64 *pTarget, FLOAT64 eta, FLOAT64 alpha, UINT32 len);
UINT32 NetTrain2(struct _sNNET *pObj, FLOAT64 *pIn, FLOAT64 *pOut, FLOAT64 *pError, FLOAT64 eta, FLOAT64 alpha, UINT32 len);
UINT32 NetRTRL(struct _sNNET *pObj, FLOAT64 *pIn, FLOAT64 *pOut, FLOAT64 *pTarget, FLOAT64 eta, FLOAT64 alpha, UINT32 nPattern, UINT32 trainInterval);
UINT32 NetWire(struct _sNNET *pObj, UINT32 nnType);
UINT32 NetFree(struct _sNNET *pObj);
void NetPrint(struct _sNNET *pObj);
UINT32 LayerInit(struct _sLAYER *pObj, struct _sUNIT *pUnit, UINT32 nUnits, UINT32 LayID);
UINT32 LayerFree(struct _sLAYER *pObj);
UINT32 UnitInit(struct _sUNIT *pObj, UINT32 type, UINT32 FType, UINT32 ID);
UINT32 UnitSetFunc(struct _sUNIT *pObj, UINT32 FuncType);
UINT32 UnitFree(struct _sUNIT *pObj);
/* Wrapper functions */
void* nnMalloc(void *pBuffer, UINT32 size);
void* nnFree(void *pBuffer);
/* Activation Functions */
FLOAT64 Func(FLOAT64 x);
FLOAT64 dFunc(FLOAT64 x);
FLOAT64 Linear(FLOAT64 x);
FLOAT64 dLinear(FLOAT64 x);
FLOAT64 Tanh2(FLOAT64 x);
FLOAT64 dTanh2(FLOAT64 x);
FLOAT64 ASigmoid(FLOAT64 x);
FLOAT64 SSigmoid(FLOAT64 x);
FLOAT64 dSSigmoid(FLOAT64 x);
FLOAT64 Sigmoid(FLOAT64 x);
FLOAT64 dSigmoid(FLOAT64 x);
FLOAT64 FSigmoid(FLOAT64 x);
FLOAT64 dFSigmoid(FLOAT64 x);
FLOAT64 dTanh(FLOAT64 x);
FLOAT64 NullFunc(FLOAT64 x);
#endif