clean up code and loose ends on data provider changes

This commit is contained in:
Andy Eschbacher
2017-01-04 14:40:37 -05:00
parent fd9d08dfbc
commit dbee19723e
3 changed files with 14 additions and 15 deletions

View File

@@ -1,11 +1,13 @@
CREATE OR REPLACE FUNCTION
CDB_GWR(subquery text, dep_var text, ind_vars text[],
bw numeric default null, fixed boolean default False, kernel text default 'bisquare')
RETURNS table(coeffs JSON, stand_errs JSON, t_vals JSON, filtered_t_vals JSON, predicted numeric, residuals numeric, r_squared numeric, rowid bigint, bandwidth numeric)
RETURNS table(coeffs JSON, stand_errs JSON, t_vals JSON, predicted numeric, residuals numeric, r_squared numeric, rowid bigint, bandwidth numeric)
AS $$
from crankshaft.regression import gwr_cs
from crankshaft.regression import GWR
return gwr_cs.gwr(subquery, dep_var, ind_vars, bw, fixed, kernel)
gwr = GWR()
return gwr.gwr(subquery, dep_var, ind_vars, bw, fixed, kernel)
$$ LANGUAGE plpythonu;

View File

@@ -1,2 +1,3 @@
from crankshaft.regression.gwr import *
from crankshaft.regression.glm import *
from crankshaft.regression.gwr_cs import *

View File

@@ -2,20 +2,18 @@
Geographically weighted regression
"""
import numpy as np
from gwr.base.gwr import GWR
from gwr.base.gwr import GWR as pysal_GWR
from gwr.base.sel_bw import Sel_BW
import plpy
import crankshaft.pysal_utils as pu
import json
from crankshaft.analysis_data_provider import AnalysisDataProvider
class GWR:
def __init__(self, analysis_provider=None):
if analysis_provider:
self.analysis_provider = analysis_provider
def __init__(self, data_provider=None):
if data_provider:
self.data_provider = data_provider
else:
self.analysis_provider = AnalysisDataProvider()
self.data_provider = AnalysisDataProvider()
def gwr(self, subquery, dep_var, ind_vars,
bw=None, fixed=False, kernel='bisquare',
@@ -36,7 +34,7 @@ class GWR:
'ind_vars': ind_vars}
# retrieve data
query_result = self.analysis_data_provider.get_gwr(params)
query_result = self.data_provider.get_gwr(params)
# unique ids and variable names list
rowid = np.array(query_result[0]['rowid'], dtype=np.int)
@@ -64,13 +62,11 @@ class GWR:
ind_vars.insert(0, 'intercept')
# calculate bandwidth if none is supplied
plpy.notice(str(bw))
if bw is None:
bw = Sel_BW(coords, Y, X,
fixed=fixed, kernel=kernel).search()
plpy.notice(str(bw))
model = GWR(coords, Y, X, bw,
fixed=fixed, kernel=kernel).fit()
model = pysal_GWR(coords, Y, X, bw,
fixed=fixed, kernel=kernel).fit()
# containers for outputs
coeffs = []