lolamontes69

Ch12 (Neural Network dict): Programming Collective Intel

Sep 29th, 2013
87
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Python 8.64 KB | None | 0 0
  1. """ Chapter 12 Programming Collective Intelligence: Neural networks.
  2.  
  3.    As the chapter has no exercises in it I thought I would get to know the neural network
  4.    a bit better by rewriting it to use saved dictionaries instead of a sql database.
  5.    It works the same way outputting three files:
  6.       "hiddennode.json", "wordhidden.json" and "hiddenurl.json."
  7.    These files are loaded/created on initialization, and updated during training etc.
  8.    I could have probably saved the dictionaries as pickles but I felt like json :)
  9. """
  10.  
  11. from math import tanh
  12. import json
  13.  
  14. def dtanh(y):
  15.     return 1.0-y*y
  16.  
  17. class searchnet:
  18.     def __init__(self):
  19.         try:
  20.             f = open('hiddennode.json', 'r')
  21.             self.hiddennode={}
  22.             for line in f:
  23.                 self.hiddennode = json.loads(line)           # json saves keys as strings so...
  24.             for a in self.hiddennode.keys():
  25.                 self.hiddennode[int(a)]=self.hiddennode[a]   # ...for keeping track of hiddenid.
  26.                 del self.hiddennode[a]
  27.             f = open('wordhidden.json', 'r')
  28.             self.wordhidden={}
  29.             for line in f:
  30.                 self.wordhidden = json.loads(line)
  31.             f = open('hiddenurl.json', 'r')
  32.             self.hiddenurl={}
  33.             for line in f:
  34.                 self.hiddenurl = json.loads(line)
  35.         except:
  36.             print "The required dictionaries don't exist in this folder."
  37.             shall=str(raw_input('Shall I create them for you? (y/n) >'))
  38.             if shall=='y' or shall=='Y':
  39.                 self.create_dics()
  40.                 print"-"*44
  41.                 print '"hiddennode.json", "wordhidden.json" and "hiddenurl.json" created.'
  42.  
  43.     def create_dics(self):
  44.         self.hiddennode={}
  45.         self.wordhidden={}
  46.         self.hiddenurl={}
  47.         f = open('hiddennode.json', 'w')
  48.         f.write(json.dumps(self.hiddennode))
  49.         f.close()
  50.         f = open('wordhidden.json', 'w')
  51.         f.write(json.dumps(self.wordhidden))
  52.         f.close()
  53.         f = open('hiddenurl.json', 'w')
  54.         f.write(json.dumps(self.hiddenurl))
  55.         f.close()
  56.  
  57.     def save_dics(self):
  58.         f = open('hiddennode.json', 'w')
  59.         f.write(json.dumps(self.hiddennode))
  60.         f.close()
  61.         f = open('wordhidden.json', 'w')
  62.         f.write(json.dumps(self.wordhidden))
  63.         f.close()
  64.         f = open('hiddenurl.json', 'w')
  65.         f.write(json.dumps(self.hiddenurl))
  66.         f.close()
  67.  
  68.     def getstrength(self,fromid1,toid1,layer):
  69.         fromid=str(fromid1)  # json saves as strings
  70.         toid=str(toid1)
  71.         if layer==0:
  72.             try: res=self.wordhidden[fromid][toid]
  73.             except: res=-0.2
  74.         else:
  75.             try: res=self.hiddenurl[fromid][toid]
  76.             except: res=0
  77.         return float(res)
  78.  
  79.     def setstrength(self,fromid1,toid1,layer,strength):
  80.         fromid=str(fromid1)  # json saves as strings
  81.         toid=str(toid1)
  82.         if layer==0:
  83.             if str(fromid) in self.wordhidden:
  84.                 self.wordhidden[fromid][toid]=strength
  85.             else:
  86.                 self.wordhidden[fromid]={}
  87.                 self.wordhidden[fromid][toid]=strength
  88.         else:
  89.             if fromid in self.hiddenurl:
  90.                 self.hiddenurl[fromid][toid]=strength
  91.             else:
  92.                 self.hiddenurl[fromid]={}
  93.                 self.hiddenurl[fromid][toid]=strength
  94.  
  95.     def generatehiddennode(self,wordids,urls):
  96.         if len(wordids)>3: return None
  97.         createkey='_'.join(sorted([str(wi) for wi in wordids]))
  98.         hiddenid=0
  99.         notin=0
  100.         if len(self.hiddennode)==0:
  101.             self.hiddennode[0]=createkey
  102.             notin=1
  103.         else:
  104.             for a in self.hiddennode:
  105.                 if self.hiddennode[a]==createkey:
  106.                     hiddenid=a
  107.                     notin=1
  108.         if notin==0:
  109.             hiddenid=int(max(self.hiddennode))+1
  110.             self.hiddennode[hiddenid]=createkey
  111.         for wordid in wordids:
  112.             self.setstrength(wordid,hiddenid,0,1.0/len(wordids))
  113.         for urlid in urls:
  114.             self.setstrength(hiddenid,urlid,1,0.1)
  115.         self.save_dics()
  116.  
  117.     def getallhiddenids(self,wordids,urlids):
  118.         l1={}
  119.         for wordid in wordids:
  120.             if str(wordid) in self.wordhidden:
  121.                 a=self.wordhidden[str(wordid)].keys()
  122.                 for row in a: l1[row]=1
  123.         for urlid in urlids:
  124.             for a in self.hiddenurl:
  125.                 if str(urlid) in self.hiddenurl[a]:
  126.                     l1[a]=1
  127.         return l1.keys()
  128.  
  129.     def setupnetwork(self,wordids,urlids):
  130.         # Value lists
  131.         self.wordids=wordids
  132.         self.hiddenids=self.getallhiddenids(wordids,urlids)
  133.         self.urlids=urlids
  134.  
  135.         # Node outputs         strengths set to default values
  136.         self.ai = [1.0]*len(self.wordids)
  137.         self.ah = [1.0]*len(self.hiddenids)
  138.         self.ao = [1.0]*len(self.urlids)
  139.  
  140.         # Create weights matrix
  141.         self.wi =[[self.getstrength(wordid,hiddenid,0) for hiddenid in self.hiddenids] for wordid in self.wordids]
  142.         self.wo =[[self.getstrength(hiddenid,urlid,1) for urlid in self.urlids] for hiddenid in self.hiddenids]
  143.  
  144.     def feedforward(self):
  145.         # the only inputs are the query words
  146.         for i in range(len(self.wordids)):
  147.             self.ai[i] = 1.0
  148.  
  149.         # hidden activations
  150.         for j in range(len(self.hiddenids)):
  151.             sum = 0.0
  152.             for i in range(len(self.wordids)):
  153.                 sum = sum + self.ai[i] * self.wi[i][j]
  154.             self.ah[j] = tanh(sum)
  155.  
  156.         # output activations
  157.         for k in range(len(self.urlids)):
  158.             sum = 0.0
  159.             for j in range(len(self.hiddenids)):
  160.                 sum = sum + self.ah[j] * self.wo[j][k]
  161.             self.ao[k] = tanh(sum)
  162.  
  163.         return self.ao[:]
  164.  
  165.     def getresult(self,wordids,urlids):
  166.         self.setupnetwork(wordids,urlids)
  167.         return self.feedforward()
  168.  
  169.     def backPropagate(self,targets, N=0.5):
  170.         # calculate errors for output
  171.         output_deltas = [0.0]*len(self.urlids)
  172.         for k in range(len(self.urlids)):
  173.             error = targets[k]-self.ao[k]
  174.             output_deltas[k] = dtanh(self.ao[k])*error
  175.  
  176.         # calculate errors for hidden layer
  177.         hidden_deltas = [0.0]*len(self.hiddenids)
  178.         for j in range(len(self.hiddenids)):
  179.             error = 0.0
  180.             for k in range(len(self.urlids)):
  181.                 error = error+output_deltas[k]*self.wo[j][k]
  182.             hidden_deltas[j] = dtanh(self.ah[j])*error
  183.  
  184.         # update output weights
  185.         for j in range(len(self.hiddenids)):
  186.             for k in range(len(self.urlids)):
  187.                 change=output_deltas[k]*self.ah[j]
  188.                 self.wo[j][k] = self.wo[j][k] + N*change
  189.  
  190.         # update input weights
  191.         for i in range(len(self.wordids)):
  192.             for j in range(len(self.hiddenids)):
  193.                 change = hidden_deltas[j]*self.ai[i]
  194.                 self.wi[i][j] = self.wi[i][j] + N*change
  195.  
  196.     def trainquery(self,wordids,urlids,selectedurl):
  197.         # generate a hiddennode if neccessary
  198.         self.generatehiddennode(wordids,urlids)
  199.  
  200.         self.setupnetwork(wordids,urlids)
  201.         self.feedforward()
  202.         targets=[0.0]*len(urlids)
  203.         targets[urlids.index(selectedurl)]=1.0
  204.         error = self.backPropagate(targets)
  205.         self.updatedatabase()
  206.  
  207.     def updatedatabase(self):
  208.         # set them to database values
  209.         for i in range(len(self.wordids)):
  210.             for j in range(len(self.hiddenids)):
  211.                 self.setstrength(self.wordids[i],self.hiddenids[j],0,self.wi[i][j])
  212.         for j in range(len(self.hiddenids)):
  213.             for k in range(len(self.urlids)):
  214.                 self.setstrength(self.hiddenids[j],self.urlids[k],1,self.wo[j][k])
  215.         self.save_dics()
  216.        
  217. """
  218. Usage: The same as nn.py but without the database references.
  219.  
  220. import nn_rewrite1 as nn
  221. wWorld,wRiver,wBank,wFrog=101,102,103,104
  222. uWorldBank,uRiver,uEarth=201,202,203
  223.  
  224. mynet=nn.searchnet()
  225.  
  226. mynet.generatehiddennode([wWorld,wRiver],[uRiver,uEarth])
  227. mynet.hiddennode
  228. mynet.hiddenurl
  229. mynet.wordhidden
  230.  
  231. mynet.trainquery([wWorld,wBank],[uWorldBank,uRiver,uEarth],uWorldBank)
  232. mynet.getresult([wWorld,wBank],[uWorldBank,uRiver,uEarth])
  233.  
  234. allurls=[uWorldBank,uRiver,uEarth]
  235. for i in range(30):
  236.    mynet.trainquery([wWorld,wBank],allurls,uWorldBank)
  237.    mynet.trainquery([wRiver,wBank],allurls,uRiver)
  238.    mynet.trainquery([wWorld],allurls,uEarth)
  239.  
  240. mynet.getresult([wWorld,wBank],[uWorldBank,uRiver,uEarth])
  241. mynet.getresult([wWorld,wFrog],[uWorldBank,uRiver,uEarth])
  242. """
Advertisement
Add Comment
Please, Sign In to add comment