Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 48 additions & 5 deletions dq0/sdk/cli/experiment.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
"""

from dq0.sdk.cli.api import routes
from dq0.sdk.cli import Data
from dq0.sdk.cli.runner import DataRunner, ModelRunner
from dq0.sdk.errors import DQ0SDKError, checkSDKResponse

Expand Down Expand Up @@ -54,6 +55,7 @@ def __init__(self, project=None, name=None):
raise ValueError('You need to set the "name" argument')
self.project = project
self.name = name
self.datasets_used = None

def get_last_model_run(self):
"""Returns the latest ModelRunner.
Expand All @@ -75,7 +77,21 @@ def get_last_data_run(self):
"""
return DataRunner(self.project)

def run(self, entry_point='train_use_dq0_makedp', args=''):
def get_dataset_uuids(self, datasets=None):
"""Returns used dataset uuids as a single comma-separated string"""
if not datasets:
datasets = self.datasets_used

data_uuids = []
for dataset in datasets:
if isinstance(dataset, Data):
data_uuids.append(dataset.uuid)
elif type(dataset) == str:
data_uuids.append(dataset)

return ','.join(data_uuids)

def run(self, args, datasets=None):
"""Starts a training run

It calls the CLI command `model train` and returns
Expand All @@ -84,26 +100,53 @@ def run(self, entry_point='train_use_dq0_makedp', args=''):
Returns:
A new instance of the ModelRunner class for the train run.
"""
if not isinstance(args, dict):
raise TypeError('args need to passed as a dict')

response = self.project._deploy()
checkSDKResponse(response)
self.project.update_commit_uuid(response['message'])

data_uuids = self.get_dataset_uuids(datasets)

if not data_uuids or not len(data_uuids):
raise DQ0SDKError('No datasets provided. Please choose which datasets to use for this run using the'
'datasets parameter or the .for_data() method')

data = {
'project_uuid': self.project.project_uuid,
'commit_uuid': self.project.commit_uuid,
'ml_project_entry_point': entry_point,
'args': args
'experiment_name': self.name,
'args': args,
'datasets': data_uuids
}

response = self.project.client.post(routes.runs.create, data=data)
checkSDKResponse(response)
print(response['message'])
print(response)
try:
job_uuid = response['message'].split(' ')[-1]
except Exception:
raise DQ0SDKError('Could not parse new commit uuid')
return ModelRunner(self.project, job_uuid)

def for_data(self, data):
"""
Specifiy which datasets are used in query.
Args:
data (:obj:`list`) list of :obj:`dq0.sdk.cli.Data` instances included in query. Alternatively, pass a single
:obj:`dq0.sdk.cli.Data` instance.

Returns:
:obj:`dq0.sdk.cli.Query` instance with set datasets
"""
if isinstance(data, Data):
data = [data]
elif not isinstance(data, list):
raise DQ0SDKError('Please provide datasets either as list of Data objects or a single Data instance')
self.datasets_used = data
return self

def preprocess(self):
"""Starts a preprocessing run

Expand All @@ -118,5 +161,5 @@ def preprocess(self):

response = self.project.post(routes.data.preprocess, id=self.project.data_source_uuid)
checkSDKResponse(response)
print(response['message'])
print(response)
return DataRunner(self.project)
2 changes: 1 addition & 1 deletion dq0/sdk/cli/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -126,7 +126,7 @@ def predict(self, test_data):
response = self.project.client.post(
routes.model.predict, uuid=self.model_uuid, data=data)
checkSDKResponse(response)
print(response['message'])
print(response)
try:
job_uuid = response['message'].split(' ')[-1]
except Exception:
Expand Down
15 changes: 10 additions & 5 deletions dq0/sdk/cli/project.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,8 @@ class Project:

Args:
name (:obj:`str`): The name of the new project
type (:obj:`str`): Type of project template to use. Can be either 'ml', 'query', 'synth' or 'estimator'.
Defaults to 'ml'.
create (bool): True to create a new project via DQ0 CLI.
Default is True.

Expand All @@ -62,10 +64,11 @@ class Project:

"""

def __init__(self, name=None, create=True):
def __init__(self, name=None, project_type='ml', create=True):
if name is None:
raise ValueError('You need to set the "name" argument')
self.name = name
self.project_type = project_type
self.commit_uuid = ''
self.datasets = []
self.experiment_name = ''
Expand Down Expand Up @@ -128,9 +131,11 @@ def _create_new(self, name):
name (:obj:`str`): The name of the new project
"""
working_dir = os.getcwd()
response = self.client.post(routes.project.create, data={'working_dir': working_dir, 'name': name})
response = self.client.post(routes.project.create, data={'working_dir': working_dir,
'name': name,
'type': self.project_type})
checkSDKResponse(response)
print(response['message'])
print(response)

# change to working directory where the new project was created
os.chdir(working_dir)
Expand Down Expand Up @@ -242,7 +247,7 @@ def attach_data_source(self, data=None, data_uuid=None, data_name=None):
else:
raise ValueError('Missing required parameter: data (Data instance) or data_uuid and data_name')

print(response['message'])
print(response)

def detach_data_source(self, data=None, data_uuid=None, data_name=None):
"""Detaches a new data source to the project.
Expand Down Expand Up @@ -270,7 +275,7 @@ def detach_data_source(self, data=None, data_uuid=None, data_name=None):
else:
raise ValueError('Missing required parameter: data (Data instance) or data_uuid and data_name')

print(response['message'])
print(response)

def get_attached_data_sources(self):
pass
Expand Down
65 changes: 49 additions & 16 deletions dq0/sdk/cli/query.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,8 @@
Copyright 2020, Gradient Zero
All rights reserved
"""
import base64

from dq0.sdk.cli import Project
from dq0.sdk.cli.api import Client, routes
from dq0.sdk.cli.data import Data
Expand Down Expand Up @@ -58,31 +60,62 @@ def get_dataset_names(self):
"""Returns used dataset names as a single comma-separated string"""
return ','.join([dataset.name for dataset in self.datasets_used])

def execute(self, query, epsilon=1.0, tau=0.0, private_column='', permissions=None, params=None):
def get_dataset_uuids(self, datasets=None):
"""Returns used dataset uuids as a single comma-separated string"""
if not datasets:
datasets = self.datasets_used

data_uuids = []
for dataset in datasets:
if isinstance(dataset, Data):
data_uuids.append(dataset.uuid)
elif type(dataset) == str:
data_uuids.append(dataset)

return ','.join(data_uuids)

def execute(self, query, args):
"""Run a query on the data sources defined by this Query instance.

Args:
query: string containing SQL
epsilon: float; Epsilon value for differential private query. Default: 1.0
tau: float; Tau threshold value for private query. Default: 0.0
private_column: string; Private column for this query. Leave empty or omit for default value from metadata.
permissions: optional; e.g. 'households<75'
params: optional; e.g. 'p1=123'
args:
entry-point: string;
epsilon: float; Epsilon value for differential private query. Default: 1.0
tau: float; Tau threshold value for private query. Default: 0.0
private_column: string; Private column for this query. Leave empty or omit for default value from metadata.
permissions: optional; e.g. 'households<75'
params: optional; e.g. 'p1=123'
Returns:
:obj:`dq0.sdk.cli.runner.QueryRunner` instance
"""
if not isinstance(args, dict):
raise TypeError('args need to passed as a dict')

response = self.project._deploy()
checkSDKResponse(response)
self.project.update_commit_uuid(response['message'])

if 'epsilon' not in args:
args['epsilon'] = '1.0'
if 'tau' not in args:
args['tau'] = '0'
if 'entry-point' not in args:
args['entry-point'] = 'execute'
args['job-type'] = 'query.run'

query_encoded = base64.b64encode(query.encode('UTF-8'))
args['query-encoded'] = query_encoded.decode('UTF-8')

if not self.datasets_used:
raise DQ0SDKError('Please specify which datasets to use for query')
raise DQ0SDKError('Please specify which datasets to use for query using the .for_data() method')
response = self.client.post(
route=routes.query.create,
data={'query': query,
'datasets': self.get_dataset_names(),
'epsilon': epsilon,
'tau': tau,
'private_column': private_column,
'permissions': permissions,
'params': params,
'project_uuid': self.project.project_uuid
route=routes.runs.create,
data={'datasets': self.get_dataset_uuids(),
'project_uuid': self.project.project_uuid,
'commit_uuid': self.project.commit_uuid,
'job-type': 'query.run',
'args': args
}
)
checkSDKResponse(response)
Expand Down
2 changes: 1 addition & 1 deletion dq0/sdk/cli/runner/runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -128,7 +128,7 @@ def _cancel(self, route, uuid):
"""
response = self.project.client.post(route, uuid=uuid)
checkSDKResponse(response)
print(response['message'])
print(response)

def wait_for_completion(self, verbose=False):
"""Loops until the state reflects the end of the run.
Expand Down
2 changes: 1 addition & 1 deletion dq0/sdk/cli/runner/state.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,7 @@ def __init__(self):

def update(self, response):
"""Updates the state representation"""
self.message = response['job_state']
self.message = response.get('job_state')
try:
self.job_uuid = response.get('job_uuid')
self.state = response.get('job_state')
Expand Down
Loading