Merge pull request #9440 from anuprulez/recommend_tools_run

Tool recommendation
This commit is contained in:
Aysam Guerler
2020-03-19 17:54:09 -04:00
committed by GitHub
14 changed files with 761 additions and 0 deletions
@@ -0,0 +1,241 @@
<template>
<div id="tool-recommendation" class="tool-recommendation-view">
<div v-if="!deprecated" class="infomessagelarge">
<h4>Tool recommendation</h4>
You have used {{ getToolId }} tool. For further analysis, you could try using the following/recommended
tools. The recommended tools are shown in the decreasing order of their scores predicted using machine
learning analysis on workflows. A tool with a higher score (closer to 100%) may fit better as the following
tool than a tool with a lower score. Please click on one of the following/recommended tools to open its
definition.
</div>
<div v-else class="warningmessagelarge">You have used {{ getToolId }} tool. {{ deprecatedMessage }}</div>
</div>
</template>
<script>
import * as d3 from "d3";
import { getAppRoot } from "onload/loadConfig";
import axios from "axios";
export default {
props: {
toolId: {
type: String,
required: true
}
},
data() {
return {
deprecated: null,
deprecatedMessage: ""
};
},
created() {
this.loadRecommendations();
},
computed: {
getToolId() {
let toolId = this.toolId || "";
if (toolId.indexOf("/") > 0) {
const toolIdSlash = toolId.split("/");
toolId = toolIdSlash[toolIdSlash.length - 2];
}
return toolId;
}
},
methods: {
loadRecommendations() {
const toolId = this.getToolId;
const url = `${getAppRoot()}api/workflows/get_tool_predictions`;
axios
.post(url, {
tool_sequence: toolId
})
.then(response => {
axios.get(`${getAppRoot()}api/datatypes/mapping`).then(responseMapping => {
const predData = response.data.predicted_data;
const datatypesMapping = responseMapping.data;
const extToType = datatypesMapping.ext_to_class_name;
const typeToType = datatypesMapping.class_to_classes;
this.deprecated = predData.is_deprecated;
if (response.data !== null && predData.children.length > 0) {
const filteredData = {};
const compatibleTools = {};
const filteredChildren = [];
const outputDatatypes = predData.o_extensions;
const children = predData.children;
for (const nameObj of children.entries()) {
const inputDatatypes = nameObj[1].i_extensions;
for (const out_t of outputDatatypes.entries()) {
for (const in_t of inputDatatypes.entries()) {
const child = extToType[out_t[1]];
const parent = extToType[in_t[1]];
if (
(typeToType[child] && parent in typeToType[child]) === true ||
out_t[1] === "input" ||
out_t[1] === "_sniff_" ||
out_t[1] === "input_collection"
) {
compatibleTools[nameObj[1].tool_id] = nameObj[1].name;
break;
}
}
}
}
for (const id in compatibleTools) {
for (const nameObj of children.entries()) {
if (nameObj[1].tool_id === id) {
filteredChildren.push(nameObj[1]);
break;
}
}
}
filteredData.o_extensions = predData.o_extensions;
filteredData.name = predData.name;
filteredData.children = filteredChildren;
if (filteredChildren.length > 0 && this.deprecated === false) {
this.renderD3Tree(filteredData);
} else if (this.deprecated === true) {
this.deprecatedMessage = predData.message;
}
}
});
});
},
renderD3Tree(predictedTools) {
const duration = 750;
const x = 620;
const y = 260;
const tree = d3.layout.tree().size([y, x]);
const diagonal = d3.svg.diagonal().projection(d => {
return [d.y, d.x];
});
const svg = d3
.select("#tool-recommendation")
.append("svg")
.attr("class", "tree-size")
.append("g")
.attr("transform", "translate(" + 250 + "," + 20 + ")");
let i = 0;
let root = null;
const update = source => {
// Compute the new tree layout.
const nodes = tree.nodes(root).reverse();
const links = tree.links(nodes);
// Normalize for fixed-depth.
nodes.forEach(d => {
d.y = d.depth * 180;
});
// Update the nodes…
const node = svg.selectAll("g.node").data(nodes, d => {
return d.id || (d.id = ++i);
});
// Enter any new nodes at the parent's previous position.
const nodeEnter = node
.enter()
.append("g")
.attr("class", "node")
.attr("transform", d => {
return "translate(" + source.y0 + "," + source.x0 + ")";
})
.on("click", click);
nodeEnter.append("circle").attr("r", 1e-6);
nodeEnter
.append("text")
.attr("x", d => {
return d.children || d._children ? -10 : 10;
})
.attr("dy", ".35em")
.attr("text-anchor", d => {
return d.children || d._children ? "end" : "start";
})
.text(d => {
return d.name;
})
.attr("class", "node-enter");
nodeEnter.append("title").text(d => {
return d.children || d._children ? "Click to collapse" : "Click to open tool definition";
});
// Transition nodes to their new position.
const nodeUpdate = node
.transition()
.duration(duration)
.attr("transform", d => {
return "translate(" + d.y + "," + d.x + ")";
});
nodeUpdate.select("circle").attr("r", 4.5);
nodeUpdate.select("text").attr("class", "node-update");
// Transition exiting nodes to the parent's new position.
const nodeExit = node
.exit()
.transition()
.duration(duration)
.attr("transform", d => {
return "translate(" + source.y + "," + source.x + ")";
})
.remove();
nodeExit.select("circle").attr("r", 1e-6);
nodeExit.select("text").attr("class", "node-enter");
// Update the links
const link = svg.selectAll("path.link").data(links, d => {
return d.target.id;
});
// Enter any new links at the parent's previous position.
link.enter()
.insert("path", "g")
.attr("class", "link")
.attr("d", d => {
const o = { x: source.x0, y: source.y0 };
return diagonal({ source: o, target: o });
});
// Transition links to their new position.
link.transition()
.duration(duration)
.attr("d", diagonal);
// Transition exiting nodes to the parent's new position.
link.exit()
.transition()
.duration(duration)
.attr("d", d => {
const o = { x: source.x, y: source.y };
return diagonal({ source: o, target: o });
})
.remove();
// Stash the old positions for transition.
nodes.forEach(d => {
d.x0 = d.x;
d.y0 = d.y;
});
};
// Toggle children on click.
const click = d => {
if (d.children) {
d._children = d.children;
d.children = null;
} else {
d.children = d._children;
d._children = null;
}
update(d);
const tId = d.tool_id;
if (tId !== undefined && tId !== "undefined" && tId !== null && tId !== "") {
document.location.href = `${getAppRoot()}tool_runner?tool_id=${tId}`;
}
};
const collapse = d => {
if (d.children) {
d._children = d.children;
d._children.forEach(collapse);
d.children = null;
}
};
root = predictedTools;
root.x0 = y / 2;
root.y0 = 0;
root.children.forEach(collapse);
update(root);
}
}
};
</script>
@@ -13,6 +13,7 @@ import ToolFormBase from "mvc/tool/tool-form-base";
import Webhooks from "mvc/webhooks";
import Vue from "vue";
import ToolEntryPoints from "components/ToolEntryPoints/ToolEntryPoints";
import ToolRecommendation from "components/ToolRecommendation";
const View = Backbone.View.extend({
initialize: function(options) {
@@ -268,6 +269,19 @@ const View = Backbone.View.extend({
}
}
this.$el.append(this._templateSuccess(response, job_def));
const enable_tool_recommendations = Galaxy.config.enable_tool_recommendations;
if (enable_tool_recommendations === true || enable_tool_recommendations === "true") {
// show tool recommendations
const ToolRecommendationInstance = Vue.extend(ToolRecommendation);
const vm = document.createElement("div");
this.$el.append(vm);
const instance = new ToolRecommendationInstance({
propsData: {
toolId: job_def.tool_id
}
});
instance.$mount(vm);
}
this.$el.parent().scrollTop(0);
// Show Webhook if job is running
if (response.jobs && response.jobs.length > 0) {
+35
View File
@@ -1650,3 +1650,38 @@ body.reports {
bottom: 0;
background: url(../../images/largespinner.gif) no-repeat center center fixed;
}
/* Used for tree in tool recommendations */
.tool-recommendation-view {
.node {
cursor: pointer;
circle {
fill: lighten($brand-primary, 20%);
stroke: lighten($brand-primary, 20%);
stroke-width: 0.3rem;
}
text {
font: 0.75rem sans-serif;
}
}
.node-enter {
fill-opacity: 1e-6;
}
.node-update {
fill-opacity: 1;
}
.tree-size {
width: 96%;
height: 50%;
position: absolute;
}
.link {
fill: none;
stroke: lighten($brand-primary, 20%);
stroke-width: 0.3rem;
}
}
@@ -0,0 +1,48 @@
# Provide a list of tools which are deprecated. These tools would be removed the list of
# recommended tools by the deep learning model. These tools will be removed from the recommendations and a warning message, set as the 'text_message',
# is shown instead of their recommendations and when they are executed.
# Format:
# <<tool id>>:
# - is_deprecate: True
# - text_message: <<message to show once they are recommended by the deep learning model>>
#
# For the following example, the tool 'cufflinks' is deprecated. It will be removed from the recommendations and a warning message, set as the 'text_message',
# is shown instead of its recommendations and when it is executed:
#
# cufflinks:
# - is_deprecated: True
# text_message: It is deprecated
#
#
#
# Provide list of tools to be recommended. These tools will either be appended to the tools recommended by deep learning model or
# completely overwrite them with these tools.
# Format:
# <<tool id>>:
# - tool_id: <<recommended tool id 1>>
# name: <<display text: name and description of the tool>>
# i_extensions:
# - <<input file extension>>
# - <<input file extension>>
# - tool_id: <<recommended tool id 2>>
# name: <<display text: name and description of the tool>>
# i_extensions:
# - <<input file extension of this tool>>
# - <<input file extension of this tool>>
# <<tool id>>:
# - tool_id ...
#
# For the following example, the tools 'cat1' and 'sort1' are shown as the recommendations for tool 'Filter1':
#
# Filter1:
# - tool_id: 'cat1'
# name: 'Concatenate datasets tail-to-head'
# i_extensions:
# - tabular
# - txt
# - tool_id: 'sort1'
# name: 'Sort data in ascending or descending order'
# i_extensions:
# - tabular
# - txt
+63
View File
@@ -4042,4 +4042,67 @@
:Type: int
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
``enable_tool_recommendations``
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
:Description:
Allow the display of tool recommendations in workflow editor and
after tool execution. If it is enabled and set to true, please
enable 'tool_recommendation_model_path' as well
:Default: ``false``
:Type: bool
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
``tool_recommendation_model_path``
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
:Description:
Set remote path of the trained model (HDF5 file) for tool
recommendation.
:Default: ``https://github.com/galaxyproject/galaxy-test-data/raw/master/tool_recommendation_model.hdf5``
:Type: str
~~~~~~~~~~~~~~~~~~~~~~~~
``topk_recommendations``
~~~~~~~~~~~~~~~~~~~~~~~~
:Description:
Set the number of predictions/recommendations to be made by the
model
:Default: ``20``
:Type: int
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
``admin_tool_recommendations_path``
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
:Description:
Set path to the additional tool preferences from Galaxy admins. It
has two blocks. One for listing deprecated tools which will be
removed from the recommendations and another is for adding
additional tools to be recommended along side those from the deep
learning model.
:Default: ``tool_recommendations_overwrite.yml``
:Type: str
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
``overwrite_model_recommendations``
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
:Description:
Overwrite or append to the tool recommendations by the deep
learning model. When set to true, all the recommendations by the
deep learning model are overwritten by the recommendations set by
an admin in a config file 'tool_recommendations_overwrite.yml'.
When set to false, the recommended tools by admins and predicted
by the deep learning model are shown.
:Default: ``false``
:Type: bool
@@ -1977,3 +1977,31 @@ galaxy:
# as threshold (above threshold: regular select fields will be used)
#select_type_workflow_threshold: -1
# Allow the display of tool recommendations in workflow editor and
# after tool execution. If it is enabled and set to true, please
# enable 'tool_recommendation_model_path' as well
#enable_tool_recommendations: false
# Set remote path of the trained model (HDF5 file) for tool
# recommendation.
#tool_recommendation_model_path: https://github.com/galaxyproject/galaxy-test-data/raw/master/tool_recommendation_model.hdf5
# Set the number of predictions/recommendations to be made by the
# model
#topk_recommendations: 20
# Set path to the additional tool preferences from Galaxy admins. It
# has two blocks. One for listing deprecated tools which will be
# removed from the recommendations and another is for adding
# additional tools to be recommended along side those from the deep
# learning model.
#admin_tool_recommendations_path: tool_recommendations_overwrite.yml
# Overwrite or append to the tool recommendations by the deep learning
# model. When set to true, all the recommendations by the deep
# learning model are overwritten by the recommendations set by an
# admin in a config file 'tool_recommendations_overwrite.yml'. When
# set to false, the recommended tools by admins and predicted by the
# deep learning model are shown.
#overwrite_model_recommendations: false
+6
View File
@@ -194,6 +194,12 @@ class ConditionalDependencies(object):
def check_influxdb(self):
return 'influxdb' in self.error_report_modules
def check_keras(self):
return asbool(self.config["enable_tool_recommendations"])
def check_tensorflow(self):
return asbool(self.config["enable_tool_recommendations"])
def optional(config_file=None):
if not config_file:
@@ -28,3 +28,7 @@ watchdog
python-gitlab
pygithub
influxdb
# Deep learning packages for tool recommendation
keras==2.2.4
tensorflow==1.12.2
+5
View File
@@ -62,6 +62,11 @@ class ConfigSerializer(base.ModelSerializer):
'communication_server_port' : _use_config,
'communication_server_host' : _use_config,
'persistent_communication_rooms' : _use_config,
'enable_tool_recommendations' : _use_config,
'tool_recommendation_model_path' : _use_config,
'admin_tool_recommendations_path' : _use_config,
'overwrite_model_recommendations' : _use_config,
'topk_recommendations' : _use_config,
'allow_user_impersonation' : _use_config,
'allow_user_creation' : _defaults_to(False), # schema default is True
'use_remote_user' : _defaults_to(None), # schema default is False; or config.single_user
+219
View File
@@ -0,0 +1,219 @@
""" Compute tool recommendations """
import json
import logging
import os
import h5py
import numpy as np
import requests
import yaml
from galaxy.tools.parameters import populate_state
from galaxy.tools.parameters.basic import workflow_building_modes
from galaxy.workflow.modules import module_factory
log = logging.getLogger(__name__)
class ToolRecommendations():
def __init__(self):
self.tool_recommendation_model_path = None
self.admin_tool_recommendations_path = None
self.deprecated_tools = dict()
self.admin_recommendations = dict()
def set_model(self, trans, remote_model_url):
"""
Create model and associated dictionaries for recommendations
"""
if not self.tool_recommendation_model_path:
# import moves from the top of file: in case the tool recommendation feature is disabled,
# keras is not downloaded because of conditional requirement and Galaxy does not build
try:
from keras.models import model_from_json
except Exception:
trans.response.status = 400
return False
self.tool_recommendation_model_path = self.download_model(remote_model_url)
self.all_tools = dict()
model_weights = list()
counter_layer_weights = 0
# collect ids and names of all the installed tools
for tool_id, tool in trans.app.toolbox.tools():
t_id_renamed = tool_id
if t_id_renamed.find("/") > -1:
t_id_renamed = t_id_renamed.split("/")[-2]
self.all_tools[t_id_renamed] = (tool_id, tool.name)
# read the hdf5 attributes
trained_model = h5py.File(self.tool_recommendation_model_path, 'r')
model_config = json.loads(trained_model['model_config'][()])
self.loaded_model = model_from_json(model_config)
# set the dictionary of tools
self.model_data_dictionary = json.loads(trained_model['data_dictionary'][()])
self.reverse_dictionary = dict((v, k) for k, v in self.model_data_dictionary.items())
# set the list of compatible tools
self.compatible_tools = json.loads(trained_model['compatible_tools'][()])
self.tool_weights = json.loads(trained_model['class_weights'][()])
self.tool_weights_sorted = dict()
# sort the tools' usage dictionary
tool_pos_sorted = [int(key) for key in self.tool_weights.keys()]
for k in tool_pos_sorted:
self.tool_weights_sorted[k] = self.tool_weights[str(k)]
# iterate through all the attributes of the model to find weights of neural network layers
for item in trained_model.keys():
if "weight_" in item:
d_key = "weight_" + str(counter_layer_weights)
weights = trained_model[d_key][()]
model_weights.append(weights)
counter_layer_weights += 1
# set the model weights
self.loaded_model.set_weights(model_weights)
return True
def collect_admin_preferences(self, admin_path):
"""
Collect preferences for recommendations of tools
set by admins as dictionaries of deprecated tools and
additional recommendations
"""
if not self.admin_tool_recommendations_path and admin_path is not None:
self.admin_tool_recommendations_path = os.path.join(os.getcwd(), admin_path)
if os.path.exists(self.admin_tool_recommendations_path):
with open(self.admin_tool_recommendations_path) as admin_recommendations:
admin_recommendation_preferences = yaml.safe_load(admin_recommendations)
if admin_recommendation_preferences:
for tool_id in admin_recommendation_preferences:
tool_info = admin_recommendation_preferences[tool_id]
if 'is_deprecated' in tool_info[0]:
self.deprecated_tools[tool_id] = tool_info[0]["text_message"]
else:
if tool_id not in self.admin_recommendations:
self.admin_recommendations[tool_id] = tool_info
def download_model(self, model_url, download_local='database/'):
"""
Download the model from remote server
"""
local_dir = os.path.join(os.getcwd(), download_local, 'tool_recommendation_model.hdf5')
# read model from remote
model_binary = requests.get(model_url)
# save model to a local directory
with open(local_dir, 'wb') as model_file:
model_file.write(model_binary.content)
return local_dir
def get_tool_extensions(self, trans, tool_id):
"""
Get the input and output extensions of a tool
"""
payload = {'type': 'tool', 'tool_id': tool_id, '_': 'true'}
inputs = payload.get('inputs', {})
trans.workflow_building_mode = workflow_building_modes.ENABLED
module = module_factory.from_dict(trans, payload)
if 'tool_state' not in payload:
module_state = {}
populate_state(trans, module.get_inputs(), inputs, module_state, check=False)
module.recover_state(module_state)
inputs = module.get_all_inputs(connectable_only=True)
outputs = module.get_all_outputs()
input_extensions = list()
output_extensions = list()
for i_ext in inputs:
input_extensions.extend(i_ext['extensions'])
for o_ext in outputs:
output_extensions.extend(o_ext['extensions'])
return input_extensions, output_extensions
def filter_tool_predictions(self, trans, prediction_data, tool_ids, tool_scores, last_tool_name):
"""
Filter tool predictions based on datatype compatibility and tool connections.
Add admin preferences to recommendations.
"""
last_compatible_tools = list()
if last_tool_name in self.compatible_tools:
last_compatible_tools = self.compatible_tools[last_tool_name].split(",")
prediction_data["is_deprecated"] = False
t_ids_scores = zip(tool_ids, tool_scores)
# form the payload of the predicted tools to be shown
for child, score in t_ids_scores:
c_dict = dict()
for t_id in self.all_tools:
# select the name and tool id if it is installed in Galaxy
if t_id == child and score > 0.0 and child in last_compatible_tools and child not in self.deprecated_tools:
full_tool_id = self.all_tools[t_id][0]
pred_input_extensions, _ = self.get_tool_extensions(trans, full_tool_id)
c_dict["name"] = self.all_tools[t_id][1] + " (" + str(score) + "%)"
c_dict["tool_id"] = full_tool_id
c_dict["i_extensions"] = list(set(pred_input_extensions))
prediction_data["children"].append(c_dict)
break
# incorporate preferences set by admins
if self.admin_tool_recommendations_path:
# filter out deprecated tools
t_ids_scores = [(tid, score) for tid, score in zip(tool_ids, tool_scores) if tid not in self.deprecated_tools]
# set the property if the last tool of the sequence is deprecated
if last_tool_name in self.deprecated_tools:
prediction_data["is_deprecated"] = True
prediction_data["message"] = self.deprecated_tools[last_tool_name]
# add the recommendations given by admins
for tool_id in self.admin_recommendations:
if last_tool_name == tool_id:
admin_recommendations = self.admin_recommendations[tool_id]
if trans.app.config.overwrite_model_recommendations is True:
prediction_data["children"] = admin_recommendations
else:
prediction_data["children"].extend(admin_recommendations)
break
# get the root name for displaying after tool run
for t_id in self.all_tools:
if t_id == last_tool_name:
prediction_data["name"] = self.all_tools[t_id][1]
break
return prediction_data
def compute_tool_prediction(self, trans, tool_sequence):
"""
Compute the predicted tools for a tool sequences
Return a payload with the tool sequences and recommended tools
Return an empty payload with just the tool sequence if anything goes wrong within the try block
"""
max_seq_len = 25
topk = trans.app.config.topk_recommendations
prediction_data = dict()
tool_sequence = tool_sequence.split(",")[::-1]
prediction_data["name"] = ",".join(tool_sequence)
prediction_data["children"] = list()
last_tool_name = tool_sequence[-1]
# do prediction only if the last is present in the collections of tools
if last_tool_name in self.model_data_dictionary:
sample = np.zeros(max_seq_len)
# get the list of datatype extensions of the last tool of the tool sequence
_, last_output_extensions = self.get_tool_extensions(trans, self.all_tools[last_tool_name][0])
prediction_data["o_extensions"] = list(set(last_output_extensions))
# get tool names without slashes and create a sequence vector
for idx, tool_name in enumerate(tool_sequence):
if tool_name.find("/") > -1:
tool_name = tool_name.split("/")[-2]
sample[idx] = int(self.model_data_dictionary[tool_name])
sample = np.reshape(sample, (1, max_seq_len))
# predict next tools for a test path
prediction = self.loaded_model.predict(sample, verbose=0)
prediction = np.reshape(prediction, (prediction.shape[1],))
# boost the predicted scores using tools' usage
weight_values = list(self.tool_weights_sorted.values())
prediction = prediction * weight_values
# normalize the predicted scores with max and sort the predictions
max_prediction = float(np.max(prediction))
if max_prediction == 0.0:
max_prediction = 1.0
prediction = prediction / max_prediction
prediction_pos = np.argsort(prediction, axis=-1)
# get topk prediction
topk_prediction_pos = prediction_pos[-topk:]
# read tool names using reverse dictionary
pred_tool_ids = [self.reverse_dictionary[int(tool_pos)] for tool_pos in topk_prediction_pos]
predicted_scores = [int(prediction[pos] * 100) for pos in topk_prediction_pos]
prediction_data = self.filter_tool_predictions(trans, prediction_data, pred_tool_ids[::-1], predicted_scores[::-1], last_tool_name)
return prediction_data
@@ -26,6 +26,7 @@ from galaxy.managers import (
from galaxy.managers.jobs import fetch_job_states, invocation_job_source_iter
from galaxy.model.item_attrs import UsesAnnotations
from galaxy.tool_shed.galaxy_install.install_manager import InstallRepositoryManager
from galaxy.tools import recommendations
from galaxy.tools.parameters import populate_state
from galaxy.tools.parameters.basic import workflow_building_modes
from galaxy.util.sanitize_html import sanitize_html
@@ -57,6 +58,7 @@ class WorkflowsAPIController(BaseAPIController, UsesStoredWorkflowMixin, UsesAnn
self.history_manager = histories.HistoryManager(app)
self.workflow_manager = workflows.WorkflowsManager(app)
self.workflow_contents_manager = workflows.WorkflowContentsManager(app)
self.tool_recommendations = recommendations.ToolRecommendations()
def __get_full_shed_url(self, url):
for name, shed_url in self.app.tool_shed_registry.tool_sheds.items():
@@ -631,6 +633,37 @@ class WorkflowsAPIController(BaseAPIController, UsesStoredWorkflowMixin, UsesAnn
'post_job_actions' : module.get_post_job_actions(inputs)
}
@expose_api
def get_tool_predictions(self, trans, payload, **kwd):
"""
POST /api/workflows/get_tool_predictions
Fetch predicted tools for a workflow
:type payload: dict
:param payload: a dictionary containing two parameters:
'tool_sequence' - comma separated sequence of tool ids
'remote_model_url' - (optional) path to the deep learning model
"""
remote_model_url = payload.get('remote_model_url', None)
if remote_model_url is None:
remote_model_url = trans.app.config.tool_recommendation_model_path
if 'tool_sequence' not in payload or remote_model_url is None:
return
is_model_set = True
recommended_tools = dict()
tool_sequence = ""
# collect tool recommendation preferences if set by admin
self.tool_recommendations.collect_admin_preferences(trans.app.config.admin_tool_recommendations_path)
# recreate the neural network model to be used for prediction
is_model_set = self.tool_recommendations.set_model(trans, remote_model_url)
if is_model_set is True:
# get the recommended tools for a tool sequence
tool_sequence = payload.get('tool_sequence', "")
recommended_tools = self.tool_recommendations.compute_tool_prediction(trans, tool_sequence)
return {
"current_tool": tool_sequence,
"predicted_data": recommended_tools
}
#
# -- Helper methods --
#
+1
View File
@@ -404,6 +404,7 @@ def populate_api_routes(webapp, app):
webapp.mapper.connect('/api/container_resolvers/{index}/toolbox', action="resolve_toolbox", controller="container_resolution", conditions=dict(method=["GET"]))
webapp.mapper.connect('/api/container_resolvers/{index}/resolve/install', action="resolve_with_install", controller="container_resolution", conditions=dict(method=["POST"]))
webapp.mapper.connect('/api/container_resolvers/{index}/toolbox/install', action="resolve_toolbox_with_install", controller="container_resolution", conditions=dict(method=["POST"]))
webapp.mapper.connect('/api/workflows/get_tool_predictions', action='get_tool_predictions', controller="workflows", conditions=dict(method=["POST"]))
webapp.mapper.resource_with_deleted('user', 'users', path_prefix='/api')
webapp.mapper.resource('genome', 'genomes', path_prefix='/api')
@@ -2982,3 +2982,44 @@ mapping:
use 0 in order to always use select2 fields,
use -1 (default) in order to always use the regular select fields,
use any other positive number as threshold (above threshold: regular select fields will be used)
enable_tool_recommendations:
type: bool
default: false
required: false
desc: |
Allow the display of tool recommendations in workflow editor and after tool execution.
If it is enabled and set to true, please enable 'tool_recommendation_model_path' as well
tool_recommendation_model_path:
type: str
default: 'https://github.com/galaxyproject/galaxy-test-data/raw/master/tool_recommendation_model.hdf5'
required: false
desc: |
Set remote path of the trained model (HDF5 file) for tool recommendation.
topk_recommendations:
type: int
default: 20
required: false
desc: |
Set the number of predictions/recommendations to be made by the model
admin_tool_recommendations_path:
type: str
required: false
default: 'tool_recommendations_overwrite.yml'
path_resolves_to: config_dir
desc: |
Set path to the additional tool preferences from Galaxy admins.
It has two blocks. One for listing deprecated tools which will be removed from the recommendations and
another is for adding additional tools to be recommended along side those from the deep learning model.
overwrite_model_recommendations:
type: bool
default: false
required: false
desc: |
Overwrite or append to the tool recommendations by the deep learning model. When set to true, all the recommendations by the deep learning model
are overwritten by the recommendations set by an admin in a config file 'tool_recommendations_overwrite.yml'. When set to false, the recommended tools
by admins and predicted by the deep learning model are shown.
+23
View File
@@ -331,6 +331,29 @@ class WorkflowsApiTestCase(BaseWorkflowsApiTestCase):
self._assert_user_has_workflow_with_name(name)
return upload_response
def test_get_tool_predictions(self):
request = {"tool_sequence": "Cut1", "remote_model_url": "https://github.com/galaxyproject/galaxy-test-data/raw/master/tool_recommendation_model.hdf5"}
actual_recommendations = ['Filter1', 'cat1', 'addValue', 'comp1', 'Grep1']
route = "workflows/get_tool_predictions"
response = self._post(route, data=request)
recommendation_response = response.json()
is_empty = bool(recommendation_response["current_tool"])
if is_empty is False:
self._assert_status_code_is(response, 400)
else:
# check Ok response from the API
self._assert_status_code_is(response, 200)
recommendation_response = response.json()
# check the input tool sequence
assert recommendation_response["current_tool"] == request["tool_sequence"]
# check non-empty predictions list
predicted_tools = recommendation_response["predicted_data"]["children"]
assert len(predicted_tools) > 0
# check for the correct predictions
for tool in predicted_tools:
assert tool["tool_id"] in actual_recommendations
break
def test_update(self):
original_workflow = self.workflow_populator.load_workflow(name="test_import")
uuids = {}