forked from yahoo/CaffeOnSpark
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathPythonApiTest.py
More file actions
64 lines (53 loc) · 2.35 KB
/
Copy pathPythonApiTest.py
File metadata and controls
64 lines (53 loc) · 2.35 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
'''
Copyright 2016 Yahoo Inc.
Licensed under the terms of the Apache 2.0 license.
Please see LICENSE file in the project root for terms.
'''
from com.yahoo.ml.caffe.CaffeOnSpark import CaffeOnSpark
from com.yahoo.ml.caffe.Config import Config
from com.yahoo.ml.caffe.DataSource import DataSource
from pyspark.sql import DataFrame
from pyspark.mllib.linalg import Vectors
from pyspark.sql import Row
from pyspark import SparkConf,SparkContext
from itertools import izip_longest
import unittest
import os.path
conf = SparkConf().setAppName("caffe-on-spark").setMaster("local[1]")
sc = SparkContext(conf=conf)
class PythonApiTest(unittest.TestCase):
def grouper(self,iterable, n, fillvalue=None):
args = [iter(iterable)] * n
return izip_longest(fillvalue=fillvalue, *args)
def setUp(self):
#Initialize all objects
self.cos=CaffeOnSpark(sc)
cmdargs = conf.get('spark.pythonargs')
self.args= dict(self.grouper(cmdargs.split(),2))
self.cfg=Config(sc,self.args)
self.train_source = DataSource(sc).getSource(self.cfg,True)
self.validation_source = DataSource(sc).getSource(self.cfg,False)
def testTrain(self):
self.cos.train(self.train_source)
self.assertTrue(os.path.isfile(self.args.get('-model').split(":")[1][3:]))
result=self.cos.features(self.validation_source)
self.assertTrue('accuracy' in result.columns)
self.assertTrue('ip1' in result.columns)
self.assertTrue('ip2' in result.columns)
self.assertTrue(result.count() > 100)
self.assertTrue(result.first()['SampleID'] == '00000000')
result=self.cos.test(self.validation_source)
self.assertTrue(result.get('accuracy') > 0.9)
def testTrainWithValidation(self):
result=self.cos.trainWithValidation(self.train_source, self.validation_source)
self.assertEqual(len(result.columns), 2)
self.assertEqual(result.columns[0], 'accuracy')
self.assertEqual(result.columns[1], 'loss')
result.show(2)
row_count = result.count()
last_row = result.rdd.zipWithIndex().filter(lambda (row,index): index==(row_count - 1)).collect()[0][0]
finalAccuracy = last_row[0][0]
self.assertTrue(finalAccuracy > 0.8)
finalLoss = last_row[1][0]
self.assertTrue(finalLoss < 0.5)
unittest.main(verbosity=2)