lolamontes69

nn.py for Collective Intelligence Chapter 4

Jun 20th, 2013
55
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Python 5.82 KB | None | 0 0
  1. from math import tanh
  2. from pysqlite2 import dbapi2 as sqlite
  3.  
  4. def dtanh(y):
  5.     return 1.0-y*y
  6.  
  7. class searchnet:
  8.     def __init__(self,dbname):
  9.         self.con=sqlite.connect(dbname)
  10.  
  11.     def __del__(self):
  12.         self.con.close()
  13.  
  14.     def maketables(self):
  15.         self.con.execute('create table hiddennode(create_key)')
  16.         self.con.execute('create table wordhidden(fromid,toid,strength)')
  17.         self.con.execute('create table hiddenurl(fromid,toid,strength)')
  18.         self.con.commit()
  19.  
  20.     def getstrength(self,fromid,toid,layer):
  21.         if layer==0: table='wordhidden'
  22.         else: table='hiddenurl'
  23.         res=self.con.execute('select strength from %s where fromid=%d and toid=%d' % (table,fromid,toid)).fetchone()
  24.         if res==None:
  25.             if layer==0: return -0.2
  26.             if layer==1: return 0
  27.         return res[0]
  28.  
  29.     def setstrength(self,fromid,toid,layer,strength):
  30.         if layer==0: table='wordhidden'
  31.         else: table='hiddenurl'
  32.         res=self.con.execute('select rowid from %s where fromid=%d and toid=%d' % (table,fromid,toid)).fetchone()
  33.         if res==None:
  34.             self.con.execute('insert into %s (fromid,toid,strength) values (%d,%d,%f)' % (table,fromid,toid,strength))
  35.         else:
  36.             rowid=res[0]
  37.             self.con.execute('update %s set strength=%f where rowid=%d' % (table,strength,rowid))
  38.  
  39.     def generatehiddennode(self,wordids,urls):
  40.         if len(wordids)>3: return None
  41.         # Check if we already created a node for this set of words
  42.         createkey='_'.join(sorted([str(wi) for wi in wordids]))
  43.         res=self.con.execute("select rowid from hiddennode where create_key='%s'" % createkey).fetchone()
  44.  
  45.         # If not create it
  46.         if res==None:
  47.             cur=self.con.execute("insert into hiddennode (create_key) values ('%s')" % createkey)
  48.             hiddenid=cur.lastrowid
  49.             # Put in some default weights
  50.             for wordid in wordids:
  51.                 self.setstrength(wordid,hiddenid,0,1.0/len(wordids))
  52.             for urlid in urls:
  53.                 self.setstrength(hiddenid,urlid,1,0.1)
  54.             self.con.commit()
  55.  
  56.     def getallhiddenids(self,wordids,urlids):
  57.         l1={}
  58.         for wordid in wordids:
  59.             cur=self.con.execute('select toid from wordhidden where fromid=%d' % wordid)
  60.             for row in cur: l1[row[0]]=1
  61.         for urlid in urlids:
  62.             cur=self.con.execute('select fromid from hiddenurl where toid=%d' % urlid)
  63.             for row in cur: l1[row[0]]=1
  64.         return l1.keys()
  65.  
  66.     def setupnetwork(self,wordids,urlids):
  67.         # Value lists
  68.         self.wordids=wordids
  69.         self.hiddenids=self.getallhiddenids(wordids,urlids)
  70.         self.urlids=urlids
  71.  
  72.         # Node outputs         strengths set to default values
  73.         self.ai = [1.0]*len(self.wordids)
  74.         self.ah = [1.0]*len(self.hiddenids)
  75.         self.ao = [1.0]*len(self.urlids)
  76.  
  77.         # Create weights matrix
  78.         self.wi =[[self.getstrength(wordid,hiddenid,0) for hiddenid in self.hiddenids] for wordid in self.wordids]
  79.         self.wo =[[self.getstrength(hiddenid,urlid,1) for urlid in self.urlids] for hiddenid in self.hiddenids]
  80.  
  81.     def feedforward(self):
  82.         # the only inputs are the query words
  83.         for i in range(len(self.wordids)):
  84.             self.ai[i] = 1.0
  85.  
  86.         # hidden activations
  87.         for j in range(len(self.hiddenids)):
  88.             sum = 0.0
  89.             for i in range(len(self.wordids)):
  90.                 sum = sum + self.ai[i] * self.wi[i][j]
  91.             self.ah[j] = tanh(sum)
  92.  
  93.         # output activations
  94.         for k in range(len(self.urlids)):
  95.             sum = 0.0
  96.             for j in range(len(self.hiddenids)):
  97.                 sum = sum + self.ah[j] * self.wo[j][k]
  98.             self.ao[k] = tanh(sum)
  99.  
  100.         return self.ao[:]
  101.  
  102.     def getresult(self,wordids,urlids):
  103.         self.setupnetwork(wordids,urlids)
  104.         return self.feedforward()
  105.  
  106.     def backPropagate(self,targets, N=0.5):
  107.         # calculate errors for output
  108.         output_deltas = [0.0]*len(self.urlids)
  109.         for k in range(len(self.urlids)):
  110.             error = targets[k]-self.ao[k]
  111.             output_deltas[k] = dtanh(self.ao[k])*error
  112.  
  113.         # calculate errors for hidden layer
  114.         hidden_deltas = [0.0]*len(self.hiddenids)
  115.         for j in range(len(self.hiddenids)):
  116.             error = 0.0
  117.             for k in range(len(self.urlids)):
  118.                 error = error+output_deltas[k]*self.wo[j][k]
  119.             hidden_deltas[j] = dtanh(self.ah[j])*error
  120.  
  121.         # update output weights
  122.         for j in range(len(self.hiddenids)):
  123.             for k in range(len(self.urlids)):
  124.                 change=output_deltas[k]*self.ah[j]
  125.                 self.wo[j][k] = self.wo[j][k] + N*change
  126.  
  127.         # update input weights
  128.         for i in range(len(self.wordids)):
  129.             for j in range(len(self.hiddenids)):
  130.                 change = hidden_deltas[j]*self.ai[i]
  131.                 self.wi[i][j] = self.wi[i][j] + N*change
  132.  
  133.     def trainquery(self,wordids,urlids,selectedurl):
  134.         # generate a hiddennode if neccessary
  135.         self.generatehiddennode(wordids,urlids)
  136.  
  137.         self.setupnetwork(wordids,urlids)
  138.         self.feedforward()
  139.         targets=[0.0]*len(urlids)
  140.         targets[urlids.index(selectedurl)]=1.0
  141.         error = self.backPropagate(targets)
  142.         self.updatedatabase()
  143.  
  144.     def updatedatabase(self):
  145.         # set them to database values
  146.         for i in range(len(self.wordids)):
  147.             for j in range(len(self.hiddenids)):
  148.                 self.setstrength(self.wordids[i],self.hiddenids[j],0,self.wi[i][j])
  149.         for j in range(len(self.hiddenids)):
  150.             for k in range(len(self.urlids)):
  151.                 self.setstrength(self.hiddenids[j],self.urlids[k],1,self.wo[j][k])
  152.         self.con.commit()
Advertisement
Add Comment
Please, Sign In to add comment