From 0fca6c3c1a4817803dccbfe8f5a3ee68c334a199 Mon Sep 17 00:00:00 2001 From: Stuart Lynn Date: Thu, 26 May 2016 12:31:58 -0400 Subject: [PATCH] inital commit of similarity functions --- src/pg/sql/80_similarity_rank.sql | 13 ++++++ .../crankshaft/similarity/__init__.py | 1 + .../crankshaft/similarity/similarity.py | 43 +++++++++++++++++++ 3 files changed, 57 insertions(+) create mode 100644 src/pg/sql/80_similarity_rank.sql create mode 100644 src/py/crankshaft/crankshaft/similarity/__init__.py create mode 100644 src/py/crankshaft/crankshaft/similarity/similarity.py diff --git a/src/pg/sql/80_similarity_rank.sql b/src/pg/sql/80_similarity_rank.sql new file mode 100644 index 0000000..6c3867f --- /dev/null +++ b/src/pg/sql/80_similarity_rank.sql @@ -0,0 +1,13 @@ +CREATE OR REPLACE FUNCTION cdb_SimilarityRank(cartodb_id numeric, query string) +returns TABLE (cartodb_id NUMERIC, similarity NUMERIC) +as $$ + from crankshaft.similarity import similarity_rank + return similarity_rank(cartodb_id, query) +$$ LANGUAGE plpythonu + +CREATE OR REPLACE FUNCTION cdb_MostSimilar(cartodb_id numeric, query string ,matches numeric) +returns TABLE (cartodb_id NUMERIC, similarity NUMERIC) +as $$ + from crankshaft.similarity import most_similar + return most_similar(matches, query) +$$ LANGUAGE plpythonu diff --git a/src/py/crankshaft/crankshaft/similarity/__init__.py b/src/py/crankshaft/crankshaft/similarity/__init__.py new file mode 100644 index 0000000..7df975c --- /dev/null +++ b/src/py/crankshaft/crankshaft/similarity/__init__.py @@ -0,0 +1 @@ +from similarity import * diff --git a/src/py/crankshaft/crankshaft/similarity/similarity.py b/src/py/crankshaft/crankshaft/similarity/similarity.py new file mode 100644 index 0000000..ac39459 --- /dev/null +++ b/src/py/crankshaft/crankshaft/similarity/similarity.py @@ -0,0 +1,43 @@ +from sklearn.neighbors import BallTree +import numpy as np +import plpy + +def similarity_rank(target_cartodb_id, query): + data = plpy.execute(query) + features, target = extract_features_target(data,target_cartodb_id) + tree = train(features) + dist, ind = tree.query(target, k=len(data)) + cartodb_ids = [ dist[ind]['cartodb_id'] for index in ind ] + return cartodb_ids, dist + +def most_similar(matches,query): + data = plpy.execute(query) + features, _ = extract_features_target(data) + tree = train(features) + results = [] + for i in features: + target = features + dist,ind = tree.query(target, k=matches) + cartodb_ids = [ dist[ind]['cartodb_id'] for index in ind ] + results.append(cartodb_ids) + return cartodb_ids, results + +def train(features): + normed_features = normalize_features(featuers) + normed_target = normalize_features([target]) + tree = BallTree(normed_features, leaf_size=2) + return tree + +def normalize_features(features): + maxes = features.max(axis=0) + mins = features.min(axis=0) + return (features - mins)/(maxes-mins) + +def extract_features_target:(data, target_cartodb_id=None): + target = None + for row in data: + data.keys().difference(['the_geom', 'the_geom_webmercator','cartodb_id']) + if data['cartodb_id'] == target_cartodb_id + target = data.values() + return np.array(data), np.array(target) +