inital commit of similarity functions
This commit is contained in:
13
src/pg/sql/80_similarity_rank.sql
Normal file
13
src/pg/sql/80_similarity_rank.sql
Normal file
@@ -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
|
||||
1
src/py/crankshaft/crankshaft/similarity/__init__.py
Normal file
1
src/py/crankshaft/crankshaft/similarity/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
from similarity import *
|
||||
43
src/py/crankshaft/crankshaft/similarity/similarity.py
Normal file
43
src/py/crankshaft/crankshaft/similarity/similarity.py
Normal file
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user