From 1995721921c9e73a9003e4397105bce6efc69276 Mon Sep 17 00:00:00 2001 From: Stuart Lynn Date: Fri, 27 May 2016 10:29:15 -0400 Subject: [PATCH] adding functions to drop columns which are all nan and fill nan values with the mean of those columns --- .../crankshaft/similarity/similarity.py | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) diff --git a/src/py/crankshaft/crankshaft/similarity/similarity.py b/src/py/crankshaft/crankshaft/similarity/similarity.py index 224ea2d..8c642b7 100644 --- a/src/py/crankshaft/crankshaft/similarity/similarity.py +++ b/src/py/crankshaft/crankshaft/similarity/similarity.py @@ -1,13 +1,29 @@ from sklearn.neighbors import BallTree +import scipy.stats as stats import numpy as np import plpy def query_to_dictionary(result): return [ dict(zip(r.keys(), r.values())) for r in result ] +def drop_all_nan_columns(data): + reutrn data[~np.isnan(data).all(axis=0)] + +def fill_missing_na(data,val=None): + inds = np.where(np.isnan(data)) + if val==None: + col_mean = stats.nanmean(data,axis=0) + data[inds]=np.take(col_mean,inds[1]) + else: + data[inds]=np.take(val, inds[1]) + return data + def similarity_rank(target_cartodb_id, query): data = query_to_dictionary(plpy.execute(query)) + features, target = extract_features_target(data,target_cartodb_id) + features = fill_missing_na(drop_all_nan_columns(features)) + normed_features, normed_target = normalize_features(features,target) tree = train(normed_features) dist, ind = tree.query(normed_target, k=len(features)) @@ -26,6 +42,7 @@ def most_similar(matches,query): results.append(cartodb_ids) return cartodb_ids, results + def train(features): tree = BallTree(features, leaf_size=2) return tree