mirror of
https://github.com/galaxyproject/galaxy.git
synced 2026-09-24 16:30:27 +08:00
Merge pull request #9440 from anuprulez/recommend_tools_run
Tool recommendation
This commit is contained in:
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 --
|
||||
#
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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 = {}
|
||||
|
||||
Reference in New Issue
Block a user