git-svn-id: http://moon:8086/svn/software/trunk/libsrc/nn@1 b431acfa-c32f-4a4a-93f1-934dc6c82436
114 lines
3.3 KiB
C
Executable File
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
|