use of org.opensearch.ml.common.parameter.MLTask in project ml-commons by opensearch-project.
the class MLTrainingTaskRunner method train.
private void train(MLTask mlTask, MLInput mlInput, ActionListener<MLTaskResponse> actionListener) {
ActionListener<MLTaskResponse> listener = ActionListener.wrap(r -> actionListener.onResponse(r), e -> {
mlStats.createCounterStatIfAbsent(failureCountStat(mlTask.getFunctionName(), ActionName.TRAIN)).increment();
mlStats.getStat(ML_TOTAL_FAILURE_COUNT).increment();
actionListener.onFailure(e);
});
try {
// run training
mlTaskManager.updateTaskState(mlTask.getTaskId(), MLTaskState.RUNNING, mlTask.isAsync());
Model model = MLEngine.train(mlInput);
mlIndicesHandler.initModelIndexIfAbsent(ActionListener.wrap(indexCreated -> {
if (!indexCreated) {
listener.onFailure(new RuntimeException("No response to create ML task index"));
return;
}
// TODO: put the user into model for backend role based access control.
MLModel mlModel = new MLModel(mlInput.getAlgorithm(), model);
try (ThreadContext.StoredContext context = client.threadPool().getThreadContext().stashContext()) {
ActionListener<IndexResponse> indexResponseListener = ActionListener.wrap(r -> {
log.info("Model data indexing done, result:{}, model id: {}", r.getResult(), r.getId());
mlStats.getStat(ML_TOTAL_MODEL_COUNT).increment();
mlStats.createCounterStatIfAbsent(modelCountStat(mlTask.getFunctionName())).increment();
String returnedTaskId = mlTask.isAsync() ? mlTask.getTaskId() : null;
MLTrainingOutput output = new MLTrainingOutput(r.getId(), returnedTaskId, MLTaskState.COMPLETED.name());
listener.onResponse(MLTaskResponse.builder().output(output).build());
}, e -> {
listener.onFailure(e);
});
IndexRequest indexRequest = new IndexRequest(ML_MODEL_INDEX);
indexRequest.source(mlModel.toXContent(XContentBuilder.builder(XContentType.JSON.xContent()), ToXContent.EMPTY_PARAMS));
indexRequest.setRefreshPolicy(WriteRequest.RefreshPolicy.IMMEDIATE);
client.index(indexRequest, ActionListener.runBefore(indexResponseListener, () -> context.restore()));
} catch (Exception e) {
log.error("Failed to save ML model", e);
listener.onFailure(e);
}
}, e -> {
log.error("Failed to init ML model index", e);
listener.onFailure(e);
}));
} catch (Exception e) {
// todo need to specify what exception
log.error("Failed to train " + mlInput.getAlgorithm(), e);
listener.onFailure(e);
}
}
use of org.opensearch.ml.common.parameter.MLTask in project ml-commons by opensearch-project.
the class MLTaskManagerTests method testClear.
public void testClear() {
MLTask task1 = MLTask.builder().taskId("1").state(MLTaskState.CREATED).build();
MLTask task2 = MLTask.builder().taskId("2").state(MLTaskState.RUNNING).build();
MLTask task3 = MLTask.builder().taskId("3").state(MLTaskState.FAILED).build();
MLTask task4 = MLTask.builder().taskId("4").state(MLTaskState.COMPLETED).build();
mlTaskManager.add(task1);
mlTaskManager.add(task2);
mlTaskManager.add(task3);
mlTaskManager.add(task4);
mlTaskManager.clear();
Assert.assertFalse(mlTaskManager.contains(task1.getTaskId()));
Assert.assertFalse(mlTaskManager.contains(task2.getTaskId()));
Assert.assertFalse(mlTaskManager.contains(task3.getTaskId()));
Assert.assertFalse(mlTaskManager.contains(task4.getTaskId()));
}
use of org.opensearch.ml.common.parameter.MLTask in project ml-commons by opensearch-project.
the class MLTaskManagerTests method testGetRunningTaskCount.
public void testGetRunningTaskCount() {
MLTask task1 = MLTask.builder().taskId("1").state(MLTaskState.CREATED).build();
MLTask task2 = MLTask.builder().taskId("2").state(MLTaskState.RUNNING).build();
MLTask task3 = MLTask.builder().taskId("3").state(MLTaskState.FAILED).build();
MLTask task4 = MLTask.builder().taskId("4").state(MLTaskState.COMPLETED).build();
mlTaskManager.add(task1);
mlTaskManager.add(task2);
mlTaskManager.add(task3);
mlTaskManager.add(task4);
Assert.assertEquals(mlTaskManager.getRunningTaskCount(), 1);
}
use of org.opensearch.ml.common.parameter.MLTask in project ml-commons by opensearch-project.
the class MLTaskTests method testWriteTo.
public void testWriteTo() throws IOException {
BytesStreamOutput output = new BytesStreamOutput();
mlTask.writeTo(output);
MLTask task2 = new MLTask(output.bytes().streamInput());
assertEquals(mlTask, task2);
}
use of org.opensearch.ml.common.parameter.MLTask in project ml-commons by opensearch-project.
the class MLTaskTests method toXContent_NullValue.
public void toXContent_NullValue() throws IOException {
XContentBuilder builder = XContentBuilder.builder(XContentType.JSON.xContent());
MLTask task = MLTask.builder().build();
task.toXContent(builder, ToXContent.EMPTY_PARAMS);
String taskContent = TestHelper.xContentBuilderToString(builder);
assertEquals("{\"is_async\":false}", taskContent);
}
Aggregations