re-implement re-embeddings (#1679)

* fix mispelling

* re-implement re-embedding

* always load base embedding to fetch counts

* format

* lint

* fix tests

* lint

* fix accept handling

* test log

* more debug

* more

* more

* more

* more

* remove logging

* logging

* jsonify

* remove debugging logs

* lint

* clean up errors a bit

* fix issue found in PR review

* PR review changes
This commit is contained in:
Bruce Martin
2020-07-30 12:31:36 -07:00
committed by GitHub
parent bd147abb3f
commit 75cb513dd9
18 changed files with 225 additions and 179 deletions
+20 -9
View File
@@ -5,6 +5,23 @@ action creators related to embeddings choice
import { AnnoMatrixObsCrossfilter } from "../annoMatrix"; import { AnnoMatrixObsCrossfilter } from "../annoMatrix";
import { _setEmbeddingSubset } from "../util/stateManager/viewStackHelpers"; import { _setEmbeddingSubset } from "../util/stateManager/viewStackHelpers";
export async function _switchEmbedding(prevAnnoMatrix, newEmbeddingName) {
/*
DRY helper used by this and reembedding action creators
*/
const base = prevAnnoMatrix.base();
const embeddingDf = await base.fetch("emb", newEmbeddingName);
const annoMatrix = _setEmbeddingSubset(prevAnnoMatrix, embeddingDf);
const obsCrossfilter = await new AnnoMatrixObsCrossfilter(annoMatrix).select(
"emb",
newEmbeddingName,
{
mode: "all",
}
);
return [annoMatrix, obsCrossfilter];
}
export const layoutChoiceAction = (newLayoutChoice) => async ( export const layoutChoiceAction = (newLayoutChoice) => async (
dispatch, dispatch,
getState getState
@@ -14,15 +31,9 @@ export const layoutChoiceAction = (newLayoutChoice) => async (
layout. layout.
*/ */
const { annoMatrix: prevAnnoMatrix } = getState(); const { annoMatrix: prevAnnoMatrix } = getState();
const [annoMatrix, obsCrossfilter] = await _switchEmbedding(
const embeddingDf = await prevAnnoMatrix.base().fetch("emb", newLayoutChoice); prevAnnoMatrix,
const annoMatrix = _setEmbeddingSubset(prevAnnoMatrix, embeddingDf); newLayoutChoice
const obsCrossfilter = await new AnnoMatrixObsCrossfilter(annoMatrix).select(
"emb",
newLayoutChoice,
{
mode: "all",
}
); );
dispatch({ dispatch({
type: "set layout choice", type: "set layout choice",
+14 -21
View File
@@ -1,10 +1,10 @@
import { API } from "../globals"; import { API } from "../globals";
import { MatrixFBS } from "../util/stateManager";
import { import {
postNetworkErrorToast, postNetworkErrorToast,
postAsyncSuccessToast, postAsyncSuccessToast,
postAsyncFailureToast, postAsyncFailureToast,
} from "../components/framework/toasters"; } from "../components/framework/toasters";
import { _switchEmbedding } from "./embedding";
function abortableFetch(request, opts, timeout = 0) { function abortableFetch(request, opts, timeout = 0) {
const controller = new AbortController(); const controller = new AbortController();
@@ -24,7 +24,7 @@ function abortableFetch(request, opts, timeout = 0) {
async function doReembedFetch(dispatch, getState) { async function doReembedFetch(dispatch, getState) {
const state = getState(); const state = getState();
let cells = state.world.obsAnnotations.rowIndex.labels(); let cells = state.annoMatrix.rowIndex.labels();
// These lines ensure that we convert any TypedArray to an Array. // These lines ensure that we convert any TypedArray to an Array.
// This is necessary because JSON.stringify() does some very strange // This is necessary because JSON.stringify() does some very strange
@@ -54,10 +54,7 @@ async function doReembedFetch(dispatch, getState) {
}); });
const res = await af.ready(); const res = await af.ready();
if ( if (res.ok && res.headers.get("Content-Type").includes("application/json")) {
res.ok &&
res.headers.get("Content-Type").includes("application/octet-stream")
) {
return res; return res;
} }
@@ -67,7 +64,6 @@ async function doReembedFetch(dispatch, getState) {
if (body && body.length > 0) { if (body && body.length > 0) {
msg = `${msg} -- ${body}`; msg = `${msg} -- ${body}`;
} }
postNetworkErrorToast(msg);
throw new Error(msg); throw new Error(msg);
} }
@@ -78,17 +74,24 @@ export function requestReembed() {
return async (dispatch, getState) => { return async (dispatch, getState) => {
try { try {
const res = await doReembedFetch(dispatch, getState); const res = await doReembedFetch(dispatch, getState);
const schema = JSON.parse(res.headers.get("CxG-Schema")); const schema = await res.json();
const buffer = await res.arrayBuffer();
const df = MatrixFBS.matrixFBSToDataframe(buffer);
dispatch({ dispatch({
type: "reembed: request completed", type: "reembed: request completed",
}); });
const { annoMatrix: prevAnnoMatrix } = getState();
const base = prevAnnoMatrix.base().addEmbedding(schema);
const [annoMatrix, obsCrossfilter] = await _switchEmbedding(
base,
schema.name
);
dispatch({ dispatch({
type: "reembed: add reembedding", type: "reembed: add reembedding",
embedding: df,
schema, schema,
annoMatrix,
obsCrossfilter,
}); });
postAsyncSuccessToast("Re-embedding has completed."); postAsyncSuccessToast("Re-embedding has completed.");
} catch (error) { } catch (error) {
dispatch({ dispatch({
@@ -103,13 +106,3 @@ export function requestReembed() {
} }
}; };
} }
/* disabled until reimplementation occurs
export function reembedResetWorldToUniverse(dispatch, getState) {
const { reembedController } = getState();
if (reembedController.pendingFetch) reembedController.pendingFetch.abort();
dispatch({
type: "reembed: clear all reembeddings",
});
}
*/
+13
View File
@@ -397,6 +397,19 @@ export default class AnnoMatrix {
_subclassResponsibility(); _subclassResponsibility();
} }
// eslint-disable-next-line class-methods-use-this, no-unused-vars -- make sure subclass implements
addEmbedding(colSchema) {
/*
Add a new obs embedding to the AnnoMatrix, with provided schema.
Returns a new annomatrix.
Typical use will be to add a re-embedding that the server has calculated.
Will throw if the column schema is invalid (eg, duplicate name).
*/
_subclassResponsibility();
}
/** /**
** Private interfaces below. ** Private interfaces below.
**/ **/
+5
View File
@@ -118,6 +118,11 @@ export default class AnnoMatrixObsCrossfilter {
return new AnnoMatrixObsCrossfilter(annoMatrix, obsCrossfilter); return new AnnoMatrixObsCrossfilter(annoMatrix, obsCrossfilter);
} }
addEmbedding(colSchema) {
const annoMatrix = this.annoMatrix.addEmbedding(colSchema);
return new AnnoMatrixObsCrossfilter(annoMatrix, this.obsCrossfilter);
}
/** /**
Selection state - API is identical to ImmutableTypedCrossfilter, as these Selection state - API is identical to ImmutableTypedCrossfilter, as these
are just wrappers to lazy create indices. are just wrappers to lazy create indices.
+45 -23
View File
@@ -6,6 +6,7 @@ import {
removeObsAnnoColumn, removeObsAnnoColumn,
addObsAnnoCategory, addObsAnnoCategory,
removeObsAnnoCategory, removeObsAnnoCategory,
addObsLayout,
} from "../util/stateManager/schemaHelpers"; } from "../util/stateManager/schemaHelpers";
import { isArrayOrTypedArray } from "../util/typeHelpers"; import { isArrayOrTypedArray } from "../util/typeHelpers";
import { _whereCacheCreate } from "./whereCache"; import { _whereCacheCreate } from "./whereCache";
@@ -47,9 +48,9 @@ export default class AnnoMatrixLoader extends AnnoMatrix {
const colSchema = _getColumnSchema(this.schema, "obs", col); const colSchema = _getColumnSchema(this.schema, "obs", col);
_writableCategoryTypeCheck(colSchema); // throws on error _writableCategoryTypeCheck(colSchema); // throws on error
const o = this._clone(); const newAnnoMatrix = this._clone();
o.schema = addObsAnnoCategory(this.schema, col, category); newAnnoMatrix.schema = addObsAnnoCategory(this.schema, col, category);
return o; return newAnnoMatrix;
} }
async removeObsAnnoCategory(col, category, unassignedCategory) { async removeObsAnnoCategory(col, category, unassignedCategory) {
@@ -59,13 +60,17 @@ export default class AnnoMatrixLoader extends AnnoMatrix {
const colSchema = _getColumnSchema(this.schema, "obs", col); const colSchema = _getColumnSchema(this.schema, "obs", col);
_writableCategoryTypeCheck(colSchema); // throws on error _writableCategoryTypeCheck(colSchema); // throws on error
const o = await this.resetObsColumnValues( const newAnnoMatrix = await this.resetObsColumnValues(
col, col,
category, category,
unassignedCategory unassignedCategory
); );
o.schema = removeObsAnnoCategory(o.schema, col, category); newAnnoMatrix.schema = removeObsAnnoCategory(
return o; newAnnoMatrix.schema,
col,
category
);
return newAnnoMatrix;
} }
dropObsColumn(col) { dropObsColumn(col) {
@@ -75,10 +80,10 @@ export default class AnnoMatrixLoader extends AnnoMatrix {
const colSchema = _getColumnSchema(this.schema, "obs", col); const colSchema = _getColumnSchema(this.schema, "obs", col);
_writableCheck(colSchema); // throws on error _writableCheck(colSchema); // throws on error
const o = this._clone(); const newAnnoMatrix = this._clone();
o._cache.obs = this._cache.obs.dropCol(col); newAnnoMatrix._cache.obs = this._cache.obs.dropCol(col);
o.schema = removeObsAnnoColumn(this.schema, col); newAnnoMatrix.schema = removeObsAnnoColumn(this.schema, col);
return o; return newAnnoMatrix;
} }
addObsColumn(colSchema, Ctor, value) { addObsColumn(colSchema, Ctor, value) {
@@ -98,7 +103,7 @@ export default class AnnoMatrixLoader extends AnnoMatrix {
throw new Error("column already exists"); throw new Error("column already exists");
} }
const o = this._clone(); const newAnnoMatrix = this._clone();
let data; let data;
if (isArrayOrTypedArray(value)) { if (isArrayOrTypedArray(value)) {
if (value.constructor !== Ctor) if (value.constructor !== Ctor)
@@ -109,10 +114,13 @@ export default class AnnoMatrixLoader extends AnnoMatrix {
} else { } else {
data = new Ctor(this.nObs).fill(value); data = new Ctor(this.nObs).fill(value);
} }
o._cache.obs = this._cache.obs.withCol(colName, data); newAnnoMatrix._cache.obs = this._cache.obs.withCol(colName, data);
_normalizeCategoricalSchema(colSchema, o._cache.obs.col(colName)); _normalizeCategoricalSchema(
o.schema = addObsAnnoColumn(this.schema, colName, colSchema); colSchema,
return o; newAnnoMatrix._cache.obs.col(colName)
);
newAnnoMatrix.schema = addObsAnnoColumn(this.schema, colName, colSchema);
return newAnnoMatrix;
} }
renameObsColumn(oldCol, newCol) { renameObsColumn(oldCol, newCol) {
@@ -155,13 +163,13 @@ export default class AnnoMatrixLoader extends AnnoMatrix {
data[idx] = value; data[idx] = value;
} }
const o = this._clone(); const newAnnoMatrix = this._clone();
o._cache.obs = this._cache.obs.replaceColData(col, data); newAnnoMatrix._cache.obs = this._cache.obs.replaceColData(col, data);
const { categories } = colSchema; const { categories } = colSchema;
if (!categories?.includes(value)) { if (!categories?.includes(value)) {
o.schema = addObsAnnoCategory(this.schema, col, value); newAnnoMatrix.schema = addObsAnnoCategory(this.schema, col, value);
} }
return o; return newAnnoMatrix;
} }
async resetObsColumnValues(col, oldValue, newValue) { async resetObsColumnValues(col, oldValue, newValue) {
@@ -185,13 +193,27 @@ export default class AnnoMatrixLoader extends AnnoMatrix {
if (data[i] === oldValue) data[i] = newValue; if (data[i] === oldValue) data[i] = newValue;
} }
const o = this._clone(); const newAnnoMatrix = this._clone();
o._cache.obs = this._cache.obs.replaceColData(col, data); newAnnoMatrix._cache.obs = this._cache.obs.replaceColData(col, data);
const { categories } = colSchema; const { categories } = colSchema;
if (!categories?.includes(newValue)) { if (!categories?.includes(newValue)) {
o.schema = addObsAnnoCategory(this.schema, col, newValue); newAnnoMatrix.schema = addObsAnnoCategory(this.schema, col, newValue);
} }
return o; return newAnnoMatrix;
}
addEmbedding(colSchema) {
/*
add new layout to the obs embeddings
*/
const { name: colName } = colSchema;
if (_getColumnSchema(this.schema, "emb", colName)) {
throw new Error("column already exists");
}
const newAnnoMatrix = this._clone();
newAnnoMatrix.schema = addObsLayout(this.schema, colSchema);
return newAnnoMatrix;
} }
/** /**
+46 -31
View File
@@ -17,59 +17,74 @@ class AnnoMatrixView extends AnnoMatrix {
} }
addObsAnnoCategory(col, category) { addObsAnnoCategory(col, category) {
const o = this._clone(); const newAnnoMatrix = this._clone();
o.viewOf = this.viewOf.addObsAnnoCategory(col, category); newAnnoMatrix.viewOf = this.viewOf.addObsAnnoCategory(col, category);
o.schema = o.viewOf.schema; newAnnoMatrix.schema = newAnnoMatrix.viewOf.schema;
return o; return newAnnoMatrix;
} }
async removeObsAnnoCategory(col, category, unassignedCategory) { async removeObsAnnoCategory(col, category, unassignedCategory) {
const o = this._clone(); const newAnnoMatrix = this._clone();
o.viewOf = await this.viewOf.removeObsAnnoCategory( newAnnoMatrix.viewOf = await this.viewOf.removeObsAnnoCategory(
col, col,
category, category,
unassignedCategory unassignedCategory
); );
o.schema = o.viewOf.schema; newAnnoMatrix.schema = newAnnoMatrix.viewOf.schema;
return o; return newAnnoMatrix;
} }
dropObsColumn(col) { dropObsColumn(col) {
const o = this._clone(); const newAnnoMatrix = this._clone();
o.viewOf = this.viewOf.dropObsColumn(col); newAnnoMatrix.viewOf = this.viewOf.dropObsColumn(col);
o._cache.obs = this._cache.obs.dropCol(col); newAnnoMatrix._cache.obs = this._cache.obs.dropCol(col);
o.schema = o.viewOf.schema; newAnnoMatrix.schema = newAnnoMatrix.viewOf.schema;
return o; return newAnnoMatrix;
} }
addObsColumn(colSchema, Ctor, value) { addObsColumn(colSchema, Ctor, value) {
const o = this._clone(); const newAnnoMatrix = this._clone();
o.viewOf = this.viewOf.addObsColumn(colSchema, Ctor, value); newAnnoMatrix.viewOf = this.viewOf.addObsColumn(colSchema, Ctor, value);
o.schema = o.viewOf.schema; newAnnoMatrix.schema = newAnnoMatrix.viewOf.schema;
return o; return newAnnoMatrix;
} }
renameObsColumn(oldCol, newCol) { renameObsColumn(oldCol, newCol) {
const o = this._clone(); const newAnnoMatrix = this._clone();
o.viewOf = this.viewOf.renameObsColumn(oldCol, newCol); newAnnoMatrix.viewOf = this.viewOf.renameObsColumn(oldCol, newCol);
o.schema = o.viewOf.schema; newAnnoMatrix.schema = newAnnoMatrix.viewOf.schema;
return o; return newAnnoMatrix;
} }
async setObsColumnValues(col, rowLabels, value) { async setObsColumnValues(col, rowLabels, value) {
const o = this._clone(); const newAnnoMatrix = this._clone();
o.viewOf = await this.viewOf.setObsColumnValues(col, rowLabels, value); newAnnoMatrix.viewOf = await this.viewOf.setObsColumnValues(
o._cache.obs = this._cache.obs.dropCol(col); col,
o.schema = o.viewOf.schema; rowLabels,
return o; value
);
newAnnoMatrix._cache.obs = this._cache.obs.dropCol(col);
newAnnoMatrix.schema = newAnnoMatrix.viewOf.schema;
return newAnnoMatrix;
} }
async resetObsColumnValues(col, oldValue, newValue) { async resetObsColumnValues(col, oldValue, newValue) {
const o = this._clone(); const newAnnoMatrix = this._clone();
o.viewOf = await this.viewOf.resetObsColumnValues(col, oldValue, newValue); newAnnoMatrix.viewOf = await this.viewOf.resetObsColumnValues(
o._cache.obs = this._cache.obs.dropCol(col); col,
o.schema = o.viewOf.schema; oldValue,
return o; newValue
);
newAnnoMatrix._cache.obs = this._cache.obs.dropCol(col);
newAnnoMatrix.schema = newAnnoMatrix.viewOf.schema;
return newAnnoMatrix;
}
addEmbedding(colSchema) {
const newAnnoMatrix = this._clone();
newAnnoMatrix.viewOf = this.viewOf.addEmbedding(colSchema);
newAnnoMatrix.schema = newAnnoMatrix.viewOf.schema;
return newAnnoMatrix;
} }
} }
+1 -1
View File
@@ -101,7 +101,7 @@ export default Embedding;
const loadAllEmbeddingCounts = async ({ annoMatrix, available }) => { const loadAllEmbeddingCounts = async ({ annoMatrix, available }) => {
const embeddings = await Promise.all( const embeddings = await Promise.all(
available.map((name) => annoMatrix.fetch("emb", name)) available.map((name) => annoMatrix.base().fetch("emb", name))
); );
return available.map((name, idx) => ({ return available.map((name, idx) => ({
embeddingName: name, embeddingName: name,
+5
View File
@@ -10,6 +10,7 @@ import InformationMenu from "./infoMenu";
import Subset from "./subset"; import Subset from "./subset";
import UndoRedoReset from "./undoRedo"; import UndoRedoReset from "./undoRedo";
import DiffexpButtons from "./diffexpButtons"; import DiffexpButtons from "./diffexpButtons";
import Reembedding from "./reembedding";
import { getEmbSubsetView } from "../../util/stateManager/viewStackHelpers"; import { getEmbSubsetView } from "../../util/stateManager/viewStackHelpers";
@connect((state) => { @connect((state) => {
@@ -49,6 +50,8 @@ import { getEmbSubsetView } from "../../util/stateManager/viewStackHelpers";
tosURL: state.config?.parameters?.["about_legal_tos"], tosURL: state.config?.parameters?.["about_legal_tos"],
privacyURL: state.config?.parameters?.["about_legal_privacy"], privacyURL: state.config?.parameters?.["about_legal_privacy"],
categoricalSelection: state.categoricalSelection, categoricalSelection: state.categoricalSelection,
enableReembedding:
state.config?.parameters?.["enable-reembedding"] ?? false,
}; };
}) })
class MenuBar extends React.PureComponent { class MenuBar extends React.PureComponent {
@@ -214,6 +217,7 @@ class MenuBar extends React.PureComponent {
colorAccessor, colorAccessor,
subsetPossible, subsetPossible,
subsetResetPossible, subsetResetPossible,
enableReembedding,
} = this.props; } = this.props;
const { pendingClipPercentiles } = this.state; const { pendingClipPercentiles } = this.state;
@@ -266,6 +270,7 @@ class MenuBar extends React.PureComponent {
this.handleClipPercentileMinValueChange this.handleClipPercentileMinValueChange
} }
/> />
{enableReembedding ? <Reembedding /> : null}
<Tooltip <Tooltip
content="When a category is colored by, show labels on the graph" content="When a category is colored by, show labels on the graph"
position="bottom" position="bottom"
@@ -0,0 +1,40 @@
import React from "react";
import { connect } from "react-redux";
import { AnchorButton, ButtonGroup, Tooltip } from "@blueprintjs/core";
import * as globals from "../../globals";
import actions from "../../actions";
import styles from "./menubar.css";
@connect((state) => ({
reembedController: state.reembedController,
annoMatrix: state.annoMatrix,
}))
class Reembedding extends React.PureComponent {
render() {
const { dispatch, annoMatrix, reembedController } = this.props;
const loading = !!reembedController?.pendingFetch;
const disabled = annoMatrix.nObs === annoMatrix.schema.dataframe.nObs;
const tipContent = disabled
? "Subset cells first, then click to recompute UMAP embedding."
: "Click to recompute UMAP embedding on the current cell subset.";
return (
<ButtonGroup className={styles.menubarButton}>
<Tooltip
content={tipContent}
position="bottom"
hoverOpenDelay={globals.tooltipHoverOpenDelay}
>
<AnchorButton
icon="new-object"
disabled={disabled}
onClick={() => dispatch(actions.requestReembed())}
loading={loading}
/>
</Tooltip>
</ButtonGroup>
);
}
}
export default Reembedding;
+1 -3
View File
@@ -18,7 +18,7 @@ import autosave from "./autosave";
import ontology from "./ontology"; import ontology from "./ontology";
import centroidLabels from "./centroidLabels"; import centroidLabels from "./centroidLabels";
import pointDialation from "./pointDilation"; import pointDialation from "./pointDilation";
import { reembedController, reembedding } from "./reembed"; import { reembedController } from "./reembed";
import { gcMiddleware as annoMatrixGC } from "../annoMatrix"; import { gcMiddleware as annoMatrixGC } from "../annoMatrix";
import undoableConfig from "./undoableConfig"; import undoableConfig from "./undoableConfig";
@@ -30,7 +30,6 @@ const Reducer = undoable(
["obsCrossfilter", obsCrossfilter], ["obsCrossfilter", obsCrossfilter],
["ontology", ontology], ["ontology", ontology],
["annotations", annotations], ["annotations", annotations],
["reembedding", reembedding],
["layoutChoice", layoutChoice], ["layoutChoice", layoutChoice],
["categoricalSelection", categoricalSelection], ["categoricalSelection", categoricalSelection],
["continuousSelection", continuousSelection], ["continuousSelection", continuousSelection],
@@ -55,7 +54,6 @@ const Reducer = undoable(
"layoutChoice", "layoutChoice",
"centroidLabels", "centroidLabels",
"annotations", "annotations",
"reembedding",
], ],
undoableConfig undoableConfig
); );
+4 -16
View File
@@ -48,27 +48,15 @@ const LayoutChoice = (
} }
case "reembed: add reembedding": { case "reembed: add reembedding": {
const { schema } = nextSharedState.annoMatrix;
const { name } = action.schema; const { name } = action.schema;
const available = Array.from(new Set(state.available).add(name)); const available = Array.from(new Set(state.available).add(name));
const currentDimNames = schema.layout.obsByName[name].dims;
return { return {
...state, ...state,
available, available,
}; current: name,
} currentDimNames,
case "reembed: clear all reembeddings": {
const { annoMatrix } = nextSharedState;
const { current } = state;
const dflt = setToDefaultLayout(annoMatrix.schema);
if (dflt.available.includes(current)) {
return {
...state,
available: dflt.available,
};
}
return {
...state,
...dflt,
}; };
} }
-35
View File
@@ -27,38 +27,3 @@ export const reembedController = (
} }
} }
}; };
/*
actual reembedding data is part of the undo/redo history
*/
export const reembedding = (
state = {
reembeddings: new Map(),
},
action
) => {
switch (action.type) {
case "reembed: add reembedding": {
const { schema, embedding } = action;
const { name } = schema.name;
const { reembeddings } = state;
return {
...state,
reembeddings: new Map(reembeddings).set(name, {
name,
schema,
embedding,
}),
};
}
case "reembed: clear all reembeddings": {
return {
...state,
reembeddings: new Map(),
};
}
default: {
return state;
}
}
};
+3 -15
View File
@@ -304,10 +304,6 @@ def layout_obs_put(request, data_adaptor):
if not data_adaptor.dataset_config.embeddings__enable_reembedding: if not data_adaptor.dataset_config.embeddings__enable_reembedding:
return abort(HTTPStatus.NOT_IMPLEMENTED) return abort(HTTPStatus.NOT_IMPLEMENTED)
preferred_mimetype = request.accept_mimetypes.best_match(["application/octet-stream"])
if preferred_mimetype != "application/octet-stream":
return abort(HTTPStatus.NOT_ACCEPTABLE)
args = request.get_json() args = request.get_json()
filter = args["filter"] if args else None filter = args["filter"] if args else None
if not filter: if not filter:
@@ -315,17 +311,9 @@ def layout_obs_put(request, data_adaptor):
method = args["method"] if args else "umap" method = args["method"] if args else "umap"
try: try:
schema, fbs = data_adaptor.compute_embedding(method, filter) schema = data_adaptor.compute_embedding(method, filter)
return make_response( return make_response(jsonify(schema), HTTPStatus.OK, {"Content-Type": "application/json"})
fbs,
HTTPStatus.OK,
{
"Content-Type": "application/octet-stream",
"CxG-Schema": json.dumps(schema),
"Access-Control-Expose-Headers": "CxG-Schema",
},
)
except NotImplementedError as e: except NotImplementedError as e:
return abort_and_log(HTTPStatus.NOT_IMPLEMENTED, str(e), include_exc_info=True) return abort_and_log(HTTPStatus.NOT_IMPLEMENTED, str(e))
except (ValueError, DisabledFeatureError, FilterError) as e: except (ValueError, DisabledFeatureError, FilterError) as e:
return abort_and_log(HTTPStatus.BAD_REQUEST, str(e), include_exc_info=True) return abort_and_log(HTTPStatus.BAD_REQUEST, str(e), include_exc_info=True)
+7 -3
View File
@@ -1,4 +1,5 @@
import importlib import importlib
import numpy as np
""" """
Wrapper for various scanpy modules. Will raise NotImplementedError if the scanpy Wrapper for various scanpy modules. Will raise NotImplementedError if the scanpy
@@ -11,8 +12,8 @@ def get_scanpy_module():
sc = importlib.import_module("scanpy") sc = importlib.import_module("scanpy")
# Future: we could enforce versions here, eg, lookat sc.__version__ # Future: we could enforce versions here, eg, lookat sc.__version__
return sc return sc
except ModuleNotFoundError: except ModuleNotFoundError as e:
raise NotImplementedError("Please install scanpy to enable UMAP re-embedding") raise NotImplementedError("Please install scanpy to enable UMAP re-embedding") from e
except Exception as e: except Exception as e:
# will capture other ImportError corner cases # will capture other ImportError corner cases
raise NotImplementedError() from e raise NotImplementedError() from e
@@ -46,4 +47,7 @@ def scanpy_umap(adata, obs_mask=None, pca_options={}, neighbors_options={}, umap
sc.pp.neighbors(adata, **neighbors_options) sc.pp.neighbors(adata, **neighbors_options)
sc.tl.umap(adata, **umap_options) sc.tl.umap(adata, **umap_options)
return adata.obsm["X_umap"] umap = adata.obsm["X_umap"]
result = np.full((obs_mask.shape[0], umap.shape[1]), np.NaN)
result[obs_mask] = umap
return result
+5 -7
View File
@@ -1,7 +1,6 @@
import warnings import warnings
import numpy as np import numpy as np
import pandas as pd
from pandas.core.dtypes.dtypes import CategoricalDtype from pandas.core.dtypes.dtypes import CategoricalDtype
import anndata import anndata
from scipy import sparse from scipy import sparse
@@ -314,16 +313,15 @@ class AnndataAdaptor(DataAdaptor):
raise FilterError("Error parsing filter") raise FilterError("Error parsing filter")
with ServerTiming.time("layout.compute"): with ServerTiming.time("layout.compute"):
X_umap = scanpy_umap(self.data, obs_mask) X_umap = scanpy_umap(self.data, obs_mask)
normalized_layout = DataAdaptor.normalize_embedding(X_umap)
# Server picks reemedding name, which must not collide with any other # Server picks reemedding name, which must not collide with any other
# embedding name generated by this backed. # embedding name generated by this backend.
name = f"reembed:{method}_{datetime.now().isoformat(timespec='milliseconds')}" name = f"reembed:{method}_{datetime.now().isoformat(timespec='milliseconds')}"
dims = [f"{name}_0", f"{name}_1"] dims = [f"{name}_0", f"{name}_1"]
df = pd.DataFrame(normalized_layout, columns=dims) layout_schema = {"name": name, "type": "float32", "dims": dims}
fbs = encode_matrix_fbs(df, col_idx=df.columns, row_idx=None) self.schema["layout"]["obs"].append(layout_schema)
schema = {"name": name, "type": "float32", "dims": dims} self.data.obsm[f"X_{name}"] = X_umap
return (schema, fbs) return layout_schema
def compute_diffexp_ttest(self, maskA, maskB, top_n=None, lfc_cutoff=None): def compute_diffexp_ttest(self, maskA, maskB, top_n=None, lfc_cutoff=None):
if top_n is None: if top_n is None:
+1 -2
View File
@@ -71,8 +71,7 @@ class DataAdaptor(metaclass=ABCMeta):
@abstractmethod @abstractmethod
def compute_embedding(self, method, filter): def compute_embedding(self, method, filter):
"""compute a new embedding on the specified obs subset, and return a """compute a new embedding on the specified obs subset, and return the embedding schema. """
tuple of (schema, fbs)."""
pass pass
@abstractmethod @abstractmethod
+5 -5
View File
@@ -237,14 +237,14 @@ class AdaptorTest(unittest.TestCase):
self.data.compute_embedding("umap", filter) self.data.compute_embedding("umap", filter)
return return
(schema, fbs) = self.data.compute_embedding("umap", filter) schema = self.data.compute_embedding("umap", filter)
self.assertIsInstance(schema["name"], str) self.assertIsInstance(schema["name"], str)
name = schema["name"] name = schema["name"]
self.assertEqual(schema["type"], "float32") self.assertEqual(schema["type"], "float32")
self.assertEqual(schema["dims"], [f"{name}_0", f"{name}_1"]) self.assertEqual(schema["dims"], [f"{name}_0", f"{name}_1"])
emb = decode_fbs.decode_matrix_FBS(fbs) emb = self.data.data.obsm[f"X_{name}"]
self.assertEqual(emb["n_rows"], 100) self.assertEqual(emb.shape, (2638, 2))
self.assertEqual(emb["n_cols"], 2) self.assertTrue(np.isfinite(emb[0:100]).all())
self.assertEqual(emb["col_idx"], [f"{name}_0", f"{name}_1"]) self.assertTrue(np.isnan(emb[100:]).all())
+10 -8
View File
@@ -73,21 +73,23 @@ class EndPoints(object):
# attempt to reembed with umap over 100 cells. # attempt to reembed with umap over 100 cells.
endpoint = "layout/obs" endpoint = "layout/obs"
url = f"{self.URL_BASE}{endpoint}" url = f"{self.URL_BASE}{endpoint}"
header = {"Accept": "application/octet-stream"}
data = {} data = {}
data["filter"] = {} data["filter"] = {}
data["filter"]["obs"] = {} data["filter"]["obs"] = {}
data["filter"]["obs"]["index"] = list(range(100)) data["filter"]["obs"]["index"] = list(range(100))
data["method"] = "umap" data["method"] = "umap"
result = self.session.put(url, headers=header, json=data) result = self.session.put(url, json=data)
self.assertEqual(result.status_code, HTTPStatus.OK) self.assertEqual(result.status_code, HTTPStatus.OK)
df = decode_fbs.decode_matrix_FBS(result.content) result_data = result.json()
self.assertEqual(df["n_rows"], 100) self.assertIsInstance(result_data, dict)
self.assertEqual(df["n_cols"], 2) self.assertEqual(result_data["type"], "float32")
cols = list(df["col_idx"]) self.assertTrue(result_data["name"].startswith("reembed:umap_"))
self.assertTrue(cols[0].startswith("reembed:umap_") and cols[0].endswith("_0")) self.assertIsInstance(result_data["dims"], list)
self.assertTrue(cols[1].startswith("reembed:umap_") and cols[1].endswith("_1")) self.assertEqual(len(result_data["dims"]), 2)
dims = result_data["dims"]
self.assertTrue(dims[0].startswith("reembed:umap_") and dims[0].endswith("_0"))
self.assertTrue(dims[1].startswith("reembed:umap_") and dims[1].endswith("_1"))
def test_bad_filter(self): def test_bad_filter(self):
endpoint = "data/var" endpoint = "data/var"