experimental re-embedding (#1186)

* first cut at re-embedding route and back-end support

* update and expand config route tests

* add scanpy_umap

* add reembedding to config route parameters

* front-end support for reembedding fetch and UI

* remove unused imports

* add loading state

* save reembedding in reducer state

* improve withColsFrom

* transmit reembed schema to client; pick unique embedding names

* display embeddings

* format

* lint

* spaces, tab size 2

* lint

* test hack for smoke-test race

* back out hack sleep

* add check for backed mode

* add unit test for reembedding

* lint

* hide re-embedding CLI param from help
This commit is contained in:
Bruce Martin
2020-03-09 17:53:30 -06:00
committed by GitHub
parent b3e9719602
commit 144b19c449
25 changed files with 928 additions and 159 deletions

View File

@@ -321,7 +321,10 @@ describe("dataframe factories", () => {
test("KeyIndex", () => {
const df = new Dataframe.Dataframe(
[2, 2],
[["red", "blue"], [true, false]],
[
["red", "blue"],
[true, false]
],
null,
new Dataframe.KeyIndex(["colors", "bools"])
);
@@ -341,7 +344,10 @@ describe("dataframe factories", () => {
test("DenseInt32Index", () => {
const df = new Dataframe.Dataframe(
[2, 2],
[["red", "blue"], [true, false]],
[
["red", "blue"],
[true, false]
],
null,
new Dataframe.DenseInt32Index([74, 75])
);
@@ -363,7 +369,10 @@ describe("dataframe factories", () => {
test("DenseInt32Index promote", () => {
const df = new Dataframe.Dataframe(
[2, 2],
[["red", "blue"], [true, false]],
[
["red", "blue"],
[true, false]
],
null,
new Dataframe.DenseInt32Index([74, 75])
);
@@ -385,7 +394,10 @@ describe("dataframe factories", () => {
test("IdentityInt32Index with last", () => {
const df = new Dataframe.Dataframe(
[2, 2],
[["red", "blue"], [true, false]],
[
["red", "blue"],
[true, false]
],
null,
null
);
@@ -407,7 +419,10 @@ describe("dataframe factories", () => {
test("IdentityInt32Index promote", () => {
const df = new Dataframe.Dataframe(
[2, 2],
[["red", "blue"], [true, false]],
[
["red", "blue"],
[true, false]
],
null,
null
);
@@ -461,7 +476,11 @@ describe("dataframe factories", () => {
*/
const dfA = new Dataframe.Dataframe(
[2, 3],
[["red", "blue"], [true, false], [1, 0]],
[
["red", "blue"],
[true, false],
[1, 0]
],
null,
new Dataframe.KeyIndex(["colors", "bools", "numbers"])
);
@@ -520,13 +539,87 @@ describe("dataframe factories", () => {
expect(dfC.col("colors").asArray()).toEqual(["red", "blue"]);
expect(dfC.col("bools").asArray()).toEqual([true, false]);
});
test("column picking", () => {
const dfEmpty = Dataframe.Dataframe.empty();
const dfA = new Dataframe.Dataframe(
[2, 1],
[["red", "blue"]],
null,
new Dataframe.KeyIndex(["colors"])
);
const dfB = new Dataframe.Dataframe(
[2, 3],
[
["red", "blue"],
[true, false],
[1, 0]
],
null,
new Dataframe.KeyIndex(["colors", "bools", "numbers"])
);
const dfX = dfEmpty.withColsFrom(dfB, ["colors", "bools"]);
expect(dfX).toBeDefined();
expect(dfX.dims).toEqual([2, 2]);
expect(dfX.colIndex.keys()).toEqual(["colors", "bools"]);
expect(dfX.rowIndex).toEqual(dfB.rowIndex);
expect(dfX.icol(0).asArray()).toEqual(dfB.icol(0).asArray());
const dfY = dfA.withColsFrom(dfB, ["numbers"]);
expect(dfY).toBeDefined();
expect(dfY.dims).toEqual([2, 2]);
expect(dfY.colIndex.keys()).toEqual(["colors", "numbers"]);
expect(dfY.rowIndex).toEqual(dfA.rowIndex);
expect(dfY.icol(0).asArray()).toEqual(dfA.icol(0).asArray());
const dfZ = dfA.withColsFrom(dfEmpty, []);
expect(dfZ).toBeDefined();
expect(dfZ.dims).toEqual(dfA.dims);
expect(dfZ.colIndex.keys()).toEqual(dfA.colIndex.keys());
expect(dfZ.rowIndex).toEqual(dfA.rowIndex);
expect(dfZ.icol(0).asArray()).toEqual(dfA.icol(0).asArray());
expect(() => dfA.withColsFrom(dfB, ["bools", "colors"])).toThrow();
});
test("column aliasing", () => {
const dfA = new Dataframe.Dataframe(
[2, 1],
[["red", "blue"]],
null,
new Dataframe.KeyIndex(["colors"])
);
const dfB = new Dataframe.Dataframe(
[2, 3],
[
["red", "blue"],
[true, false],
[1, 0]
],
null,
new Dataframe.KeyIndex(["colors", "bools", "numbers"])
);
const dfX = dfA.withColsFrom(dfB, { colors: "_colors", bools: "_bools" });
expect(dfX).toBeDefined();
expect(dfX.dims).toEqual([2, 3]);
expect(dfX.colIndex.keys()).toEqual(["colors", "_colors", "_bools"]);
expect(dfX.rowIndex).toEqual(dfA.rowIndex);
expect(dfX.icol(0).asArray()).toEqual(dfA.icol(0).asArray());
expect(dfX.col("_colors").asArray()).toBe(dfB.col("colors").asArray());
});
});
describe("dropCol", () => {
test("KeyIndex", () => {
const df = new Dataframe.Dataframe(
[2, 3],
[["red", "blue"], [true, false], [1, 0]],
[
["red", "blue"],
[true, false],
[1, 0]
],
null,
new Dataframe.KeyIndex(["colors", "bools", "numbers"])
);
@@ -545,7 +638,11 @@ describe("dataframe factories", () => {
test("IdentityInt32Index drop first", () => {
const df = new Dataframe.Dataframe(
[2, 3],
[["red", "blue"], [true, false], [1, 0]],
[
["red", "blue"],
[true, false],
[1, 0]
],
null,
null
);
@@ -565,7 +662,11 @@ describe("dataframe factories", () => {
test("IdentityInt32Index drop last", () => {
const df = new Dataframe.Dataframe(
[2, 3],
[["red", "blue"], [true, false], [1, 0]],
[
["red", "blue"],
[true, false],
[1, 0]
],
null,
null
);
@@ -585,7 +686,11 @@ describe("dataframe factories", () => {
test("DenseInt32Index", () => {
const df = new Dataframe.Dataframe(
[2, 3],
[["red", "blue"], [true, false], [1, 0]],
[
["red", "blue"],
[true, false],
[1, 0]
],
null,
new Dataframe.DenseInt32Index([102, 101, 100])
);
@@ -653,7 +758,10 @@ describe("dataframe factories", () => {
test("renameCol", () => {
const dfA = new Dataframe.Dataframe(
[2, 2],
[[true, false], [1, 0]],
[
[true, false],
[1, 0]
],
null,
new Dataframe.KeyIndex(["A", "B"])
);
@@ -671,7 +779,10 @@ describe("dataframe col", () => {
beforeEach(() => {
df = new Dataframe.Dataframe(
[2, 2],
[[true, false], [1, 0]],
[
[true, false],
[1, 0]
],
null,
new Dataframe.KeyIndex(["A", "B"])
);

View File

@@ -7,6 +7,7 @@ import {
doBinaryRequest,
dispatchNetworkErrorMessageToUser
} from "../util/actionHelpers";
import { requestReembed, reembedResetWorldToUniverse } from "./reembed";
/*
return promise to fetch the OBS annotations we need to load. Omit anything
@@ -387,6 +388,7 @@ const requestDifferentialExpression = (set1, set2, num_genes = 10) => async (
const resetWorldToUniverse = () => (dispatch, getState) => {
const { universe } = getState();
reembedResetWorldToUniverse(dispatch, getState);
dispatch({
type: "reset World to eq Universe",
universe
@@ -451,7 +453,8 @@ export default {
requestDifferentialExpression,
requestSingleGeneExpressionCountsForColoringPOST,
requestUserDefinedGene,
requestReembed,
resetWorldToUniverse,
saveObsAnnotations,
setWorldToSelection,
setWorldToSelection
};

View File

@@ -0,0 +1,113 @@
import { API } from "../globals";
import { Universe } from "../util/stateManager";
import {
postNetworkErrorToast,
postAsyncSuccessToast,
postAsyncFailureToast
} from "../components/framework/toasters";
function abortableFetch(request, opts, timeout = 0) {
const controller = new AbortController();
const { signal } = controller;
return {
abort: () => controller.abort(),
isAborted: () => signal.aborted,
ready: () => {
if (timeout) {
setTimeout(() => controller.abort(), timeout);
}
return fetch(request, { ...opts, signal });
}
};
}
async function doReembedFetch(dispatch, getState) {
const state = getState();
let cells = state.world.obsAnnotations.rowIndex.keys();
// These lines ensure that we convert any TypedArray to an Array.
// This is necessary because JSON.stringify() does some very strange
// things with TypedArrays (they are marshalled to JSON objects, rather
// than being marshalled as a JSON array).
cells = Array.isArray(cells) ? cells : Array.from(cells);
const af = abortableFetch(
`${API.prefix}${API.version}layout/obs`,
{
method: "PUT",
headers: new Headers({
Accept: "application/octet-stream",
"Content-Type": "application/json"
}),
body: JSON.stringify({
method: "umap",
filter: { obs: { index: cells } }
}),
credentials: "include"
},
60000 // 1 minute timeout
);
dispatch({
type: "reembed: request start",
abortableFetch: af
});
const res = await af.ready();
if (
res.ok &&
res.headers.get("Content-Type").includes("application/octet-stream")
) {
return res;
}
// else an error
let msg = `Unexpected HTTP response ${res.status}, ${res.statusText}`;
const body = await res.text();
if (body && body.length > 0) {
msg = `${msg} -- ${body}`;
}
postNetworkErrorToast(msg);
throw new Error(msg);
}
/*
functions below are dispatch-able
*/
export function requestReembed() {
return async (dispatch, getState) => {
try {
const res = await doReembedFetch(dispatch, getState);
const schema = JSON.parse(res.headers.get("CxG-Schema"));
const buffer = await res.arrayBuffer();
const df = Universe.matrixFBSToDataframe(buffer);
dispatch({
type: "reembed: request completed"
});
dispatch({
type: "reembed: add reembedding",
embedding: df,
schema
});
postAsyncSuccessToast("Re-embedding has completed.");
} catch (error) {
dispatch({
type: "reembed: request aborted"
});
if (error.name === "AbortError") {
postAsyncFailureToast("Re-embedding calculation was aborted.");
} else {
postNetworkErrorToast(`Re-embedding: ${error.message}`);
}
console.log("Reembed exception:", error, error.name, error.message);
}
};
}
export function reembedResetWorldToUniverse(dispatch, getState) {
const { reembedController } = getState();
if (reembedController.pendingFetch) reembedController.pendingFetch.abort();
dispatch({
type: "reembed: clear all reembeddings"
});
}

View File

@@ -2,7 +2,7 @@ import { Position, Toaster, Intent } from "@blueprintjs/core";
/** Singleton toaster instance. Create separate instances for different options. */
const ErrorToastTopCenter = Toaster.create({
const ToastTopCenter = Toaster.create({
className: "recipe-toaster",
position: Position.TOP
});
@@ -11,21 +11,38 @@ const ErrorToastTopCenter = Toaster.create({
A "user" error - eg, bad input
*/
export const postUserErrorToast = message =>
ErrorToastTopCenter.show({ message, intent: Intent.WARNING });
ToastTopCenter.show({ message, intent: Intent.WARNING });
/*
A toast the user must dismiss manually, because they need to act on its information,
ie., 8 bulk add genes out of 40 were bad. Manually see which ones and fix.
*/
export const keepAroundErrorToast = message =>
ErrorToastTopCenter.show({ message, timeout: 0, intent: Intent.WARNING });
ToastTopCenter.show({ message, timeout: 0, intent: Intent.WARNING });
/*
a hard network error
*/
export const postNetworkErrorToast = message =>
ErrorToastTopCenter.show({
ToastTopCenter.show({
message,
timeout: 30000,
intent: Intent.DANGER
});
/*
Async message to user
*/
export const postAsyncSuccessToast = message =>
ToastTopCenter.show({
message,
timeout: 10000,
intent: Intent.SUCCESS
});
export const postAsyncFailureToast = message =>
ToastTopCenter.show({
message,
timeout: 10000,
intent: Intent.WARNING
});

View File

@@ -0,0 +1,122 @@
import React from "react";
import {
AnchorButton,
ButtonGroup,
Popover,
Button,
Radio,
RadioGroup,
Tooltip,
Position
} from "@blueprintjs/core";
import { connect } from "react-redux";
import * as globals from "../../globals";
import { World } from "../../util/stateManager";
import actions from "../../actions";
@connect(state => ({
universe: state.universe,
world: state.world,
layoutChoice: state.layoutChoice,
reembedController: state.reembedController,
enableReembedding: state.config?.parameters?.["enable-reembedding"] ?? false
}))
class Embedding extends React.PureComponent {
handleLayoutChoiceChange = e => {
const { dispatch } = this.props;
dispatch({
type: "set layout choice",
layoutChoice: e.currentTarget.value
});
};
renderReembedding() {
const {
enableReembedding,
world,
universe,
dispatch,
reembedController
} = this.props;
if (!enableReembedding) return null;
const loading = !!reembedController?.pendingFetch;
const disabled = World.worldEqUniverse(world, universe);
const tipContent = disabled
? "Subset cells first, then click to recompute UMAP embedding."
: "Click to recompute UMAP embedding on the current cell subset.";
return (
<Tooltip
content={tipContent}
position="bottom"
hoverOpenDelay={globals.tooltipHoverOpenDelay}
>
<AnchorButton
icon="new-object"
style={{ marginRight: 10 }}
disabled={disabled}
onClick={() => dispatch(actions.requestReembed())}
loading={loading}
/>
</Tooltip>
);
}
render() {
const { layoutChoice } = this.props;
return (
<ButtonGroup
style={{
marginRight: 10
}}
>
<Popover
target={
<Tooltip
content="Select embedding for visualization"
position="bottom"
hoverOpenDelay={globals.tooltipHoverOpenDelay}
>
<Button
type="button"
data-testid="layout-choice"
icon="heatmap"
style={{
cursor: "pointer"
}}
/>
</Tooltip>
}
position={Position.BOTTOM_RIGHT}
content={
<div
style={{
display: "flex",
justifyContent: "flex-start",
alignItems: "flex-start",
flexDirection: "column",
padding: 10
}}
>
<RadioGroup
label="Embedding Choice"
onChange={this.handleLayoutChoiceChange}
selectedValue={layoutChoice.current}
>
{layoutChoice.available.map(name => (
<Radio label={name} value={name} key={name} />
))}
</RadioGroup>
</div>
}
/>
{this.renderReembedding()}
</ButtonGroup>
);
}
}
export default Embedding;

View File

@@ -1,20 +1,12 @@
// jshint esversion: 6
import React from "react";
import { connect } from "react-redux";
import {
Button,
ButtonGroup,
AnchorButton,
Tooltip,
Popover,
Position,
RadioGroup,
Radio,
} from "@blueprintjs/core";
import { Button, ButtonGroup, AnchorButton, Tooltip } from "@blueprintjs/core";
import * as globals from "../../globals";
import actions from "../../actions";
import CellSetButton from "./cellSetButtons";
import Clip from "./clip";
import Embedding from "./embedding";
import InformationMenu from "./infoMenu";
import Subset from "./subset";
import UndoRedoReset from "./undoRedo";
@@ -24,7 +16,6 @@ import UndoRedoReset from "./undoRedo";
world: state.world,
crossfilter: state.crossfilter,
differential: state.differential,
layoutChoice: state.layoutChoice,
graphInteractionMode: state.controls.graphInteractionMode,
clipPercentileMin: Math.round(100 * (state.world?.clipQuantiles?.min ?? 0)),
clipPercentileMax: Math.round(100 * (state.world?.clipQuantiles?.max ?? 1)),
@@ -171,14 +162,6 @@ class MenuBar extends React.Component {
this.setState({ pendingClipPercentiles: null });
};
handleLayoutChoiceChange = e => {
const { dispatch } = this.props;
dispatch({
type: "set layout choice",
layoutChoice: e.currentTarget.value
});
};
computeDiffExp = () => {
const { dispatch, differential } = this.props;
if (differential.celllist1 && differential.celllist2) {
@@ -224,7 +207,6 @@ class MenuBar extends React.Component {
return world.nObs !== universe.nObs;
};
renderDiffExp() {
/* diffexp-related buttons may be disabled */
const { disableDiffexp, differential, diffexpMayBeSlow } = this.props;
@@ -295,7 +277,6 @@ class MenuBar extends React.Component {
selectionTool,
clipPercentileMin,
clipPercentileMax,
layoutChoice,
graphInteractionMode,
aboutLink,
showCentroidLabels
@@ -335,7 +316,7 @@ class MenuBar extends React.Component {
dispatch({ type: "increment graph render counter" });
}}
/>
<ButtonGroup style={{marginRight: "10px"}}>
<ButtonGroup style={{ marginRight: "10px" }}>
<Tooltip
content={selectionTooltip}
position="bottom"
@@ -394,52 +375,7 @@ class MenuBar extends React.Component {
}}
/>
</Tooltip>
<ButtonGroup
style={{
marginRight: 10
}}
>
<Popover
target={
<Tooltip
content="Select embedding for visualization"
position="bottom"
hoverOpenDelay={globals.tooltipHoverOpenDelay}
>
<Button
type="button"
data-testid="layout-choice"
icon="heatmap"
style={{
cursor: "pointer"
}}
/>
</Tooltip>
}
position={Position.BOTTOM_RIGHT}
content={
<div
style={{
display: "flex",
justifyContent: "flex-start",
alignItems: "flex-start",
flexDirection: "column",
padding: 10
}}
>
<RadioGroup
label="Embedding Choice"
onChange={this.handleLayoutChoiceChange}
selectedValue={layoutChoice.current}
>
{layoutChoice.available.map(name => (
<Radio label={name} value={name} key={name} />
))}
</RadioGroup>
</div>
}
/>
</ButtonGroup>
<Embedding />
<Clip
pendingClipPercentiles={pendingClipPercentiles}
clipPercentileMin={clipPercentileMin}

View File

@@ -1,18 +1,17 @@
import React from "react";
import {AnchorButton, ButtonGroup, Tooltip} from "@blueprintjs/core";
import { AnchorButton, ButtonGroup, Tooltip } from "@blueprintjs/core";
import * as globals from "../../globals";
function Subset(props) {
const {
subsetPossible,
subsetResetPossible,
handleSubset,
handleSubsetReset,
handleSubsetReset
} = props;
return (
<ButtonGroup style={{marginRight: "10px"}}>
<ButtonGroup style={{ marginRight: "10px" }}>
<Tooltip
content="Subset to currently selected cells and associated metadata"
position="bottom"

View File

@@ -12,12 +12,12 @@ const ColorsReducer = (
prevSharedState
) => {
switch (action.type) {
case "universe exists, but loading is still in progress":
case "reset World to eq Universe": {
case "universe exists, but loading is still in progress": {
/* initialize everything with default colors, no mode, no color-by accessor */
const { world } = nextSharedState;
const colorMode = null;
const colorAccessor = null;
const { rgb, scale } = ColorHelpers.createColors(world, colorMode);
const { rgb, scale } = ColorHelpers.createColors(world);
return {
...state,
colorAccessor,
@@ -27,6 +27,22 @@ const ColorsReducer = (
};
}
case "reset World to eq Universe": {
/* need to rebuild colors as world may have changed, but don't switch modes */
const { world } = nextSharedState;
const { colorMode, colorAccessor } = state;
const { rgb, scale } = ColorHelpers.createColors(
world,
colorMode,
colorAccessor
);
return {
...state,
rgb,
scale
};
}
case "set clip quantiles":
case "set World to current selection": {
const { world: prevWorld, controls: prevControls } = prevSharedState;

View File

@@ -21,6 +21,7 @@ import autosave from "./autosave";
import ontology from "./ontology";
import centroidLabels from "./centroidLabels";
import pointDialation from "./pointDilation";
import { reembedController, reembedding } from "./reembed";
import undoableConfig from "./undoableConfig";
@@ -31,6 +32,7 @@ const Reducer = undoable(
["world", world],
["ontology", ontology],
["annotations", annotations],
["reembedding", reembedding],
["layoutChoice", layoutChoice],
["categoricalSelection", categoricalSelection],
["continuousSelection", continuousSelection],
@@ -42,6 +44,7 @@ const Reducer = undoable(
["responsive", responsive],
["centroidLabels", centroidLabels],
["pointDilation", pointDialation],
["reembedController", reembedController],
["autosave", autosave],
["resetCache", resetCache]
]),
@@ -57,7 +60,8 @@ const Reducer = undoable(
"differential",
"layoutChoice",
"centroidLabels",
"annotations"
"annotations",
"reembedding"
],
undoableConfig
);

View File

@@ -14,6 +14,14 @@ function bestDefaultLayout(layouts) {
return layouts[0];
}
function setToDefaultLayout(world) {
const { schema } = world;
const available = schema.layout.obs.map(v => v.name).sort();
const current = bestDefaultLayout(available);
const currentDimNames = schema.layout.obsByName[current].dims;
return { available, current, currentDimNames };
}
const LayoutChoice = (
state = {
available: [], // all available choices
@@ -24,14 +32,13 @@ const LayoutChoice = (
nextSharedState
) => {
switch (action.type) {
case "universe exists, but loading is still in progress":
case "reset World to eq Universe": {
case "universe exists, but loading is still in progress": {
// set default to default
const { schema } = nextSharedState.world;
const available = schema.layout.obs.map(v => v.name).sort();
const current = bestDefaultLayout(available);
const currentDimNames = schema.layout.obsByName[current].dims;
return { available, current, currentDimNames };
const { universe } = nextSharedState;
return {
...state,
...setToDefaultLayout(universe)
};
}
case "set layout choice": {
@@ -41,6 +48,31 @@ const LayoutChoice = (
return { ...state, current, currentDimNames };
}
case "reembed: add reembedding": {
const name = action.schema.name;
const available = Array.from(new Set(state.available).add(name));
return {
...state,
available
};
}
case "reembed: clear all reembeddings": {
const { universe } = nextSharedState;
const { current } = state;
const dflt = setToDefaultLayout(universe);
if (dflt.available.includes(current)) {
return {
...state,
available: dflt.available
};
}
return {
...state,
...dflt
};
}
default: {
return state;
}

View File

@@ -0,0 +1,64 @@
/*
controller state is not part of the undo/redo history
*/
export const reembedController = (
state = {
pendingFetch: null
},
action
) => {
switch (action.type) {
case "reembed: request start": {
return {
...state,
pendingFetch: action.abortableFetch
};
}
case "reembed: request aborted":
case "reembed: request cancel":
case "reembed: request completed": {
return {
...state,
pendingFetch: null
};
}
default: {
return state;
}
}
};
/*
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;
}
}
};

View File

@@ -4,6 +4,10 @@ import {
ControlsHelpers,
AnnotationsHelpers
} from "../util/stateManager";
import {
addObsLayout,
removeObsLayout
} from "../util/stateManager/schemaHelpers";
import clip from "../util/clip";
import quantile from "../util/quantile";
@@ -32,6 +36,7 @@ const WorldReducer = (
}
case "universe: column load success": {
/* incremental initial data load - always assumes world == universe */
const { universe } = nextSharedState;
const { dim } = action;
return {
@@ -276,6 +281,52 @@ const WorldReducer = (
};
}
case "reembed: add reembedding": {
// new embedding loaded, which *only* affects world's layout.
// It may be new, or it may replace a previous re-embedding.
const { obsLayout: origObsLayout, schema: origSchema } = state;
const { embedding, schema: embeddingSchema } = action;
const { dims, name } = embeddingSchema;
let obsLayout = origObsLayout;
let schema = origSchema;
// alias the names the server sent us, in case they were not the same as the schema
const embedingLabels = embedding.colIndex.keys();
const labels = {
[embedingLabels[0]]: dims[0],
[embedingLabels[1]]: dims[1]
};
obsLayout = obsLayout.withColsFrom(embedding, labels);
schema = addObsLayout(schema, embeddingSchema);
return {
...state,
obsLayout,
schema
};
}
case "reembed: clear all reembeddings": {
// reembedding was cleared -- remove from layout
const { obsLayout: origObsLayout, schema: origSchema } = state;
const { reembedding } = prevSharedState;
let schema = origSchema;
let obsLayout = origObsLayout;
reembedding.reembeddings.forEach((emb, name) => {
const { dims } = emb.schema;
obsLayout = obsLayout.dropCol(dims[0]);
obsLayout = obsLayout.dropCol(dims[1]);
schema = removeObsLayout(schema, name);
});
return {
...state,
obsLayout,
schema
};
}
default: {
return state;
}

View File

@@ -348,37 +348,90 @@ class Dataframe {
);
}
withColsFrom(dataframe) {
withColsFrom(dataframe, labels) {
/*
return a new dataframe containing all columns from both `this` and the
provided of dataframe.
provided dataframe argument.
The row index from `this` will be used. All dataframes must have identical
dimensionality, and no overlapping columns labels.
Special case, if either dataframe is empty, the other is returned unchanged.
Arguments:
* dataframe: a dataframe to combine with `this`
* labels: columns to pull from `dataframe` and combine with `this`. If falsey,
all columns are used. If an array, must contain a list of labels. If an
Object or Map, the key is the columns to pull, which will be stored into the
new dataframe as the value.
Example:
newDf = df.withColsFrom(otherDf); // combines all columns from both
newDf = df.withColsFrom(otherDf, ['a']); // combines df with otherDf['a']
newDf = df.withColsFrom(otherDf, {a: 'b'}); // combines df with otherDf['a'], but calls it 'b'
*/
// resolve the source and dest label names.
let srcLabels;
let dstLabels;
if (!labels) {
// combine all columns
dstLabels = dataframe.colIndex.keys();
srcLabels = dstLabels;
} else if (Array.isArray(labels)) {
// combine subset of keys with no aliasing
dstLabels = labels;
srcLabels = labels;
} else if (labels instanceof Map) {
// aliasing with a Map
srcLabels = Array.from(labels.keys());
dstLabels = Array.from(labels.values());
} else {
// aliasing with an Object
srcLabels = Object.keys(labels);
dstLabels = Object.values(labels);
}
// if datafame is empty, and no specific labels specified, noop.
if (dataframe.isEmpty()) {
if (!labels || srcLabels.length === 0) return this;
throw new Error("Empty dataframe, unable to pick columns");
}
if (this.isEmpty()) {
// 1. subset dataframe from source keys
// 2. alias names
dataframe = dataframe.subset(null, srcLabels);
for (let i = 0; i < srcLabels.length; i += 1) {
dataframe = dataframe.renameCol(srcLabels[i], dstLabels[i]);
}
return dataframe;
}
if (dataframe.isEmpty()) {
return this;
// otherwise, bulid a new dataframe combining columns from both
const srcOffsets = srcLabels.map(l => dataframe.colIndex.getOffset(l));
// check for label collisions
if (dstLabels.some(this.hasCol, this)) {
throw new Error("duplicate key collision");
}
this.colIndex.keys().forEach(key => {
if (dataframe.has(key)) {
throw new Error("duplicate key collision");
}
});
const dims = [this.dims[0], this.dims[1] + dataframe.dims[1]];
// const dims = [this.dims[0], this.dims[1] + dataframe.dims[1]];
const dims = [this.dims[0], this.dims[1] + srcOffsets.length];
const { rowIndex } = this;
const columns = [...this.__columns, ...dataframe.__columns];
const colIndex = this.colIndex.withLabels(dataframe.colIndex.keys());
const columns = [
...this.__columns,
...srcOffsets.map(i => dataframe.__columns[i])
];
const colIndex = this.colIndex.withLabels(dstLabels);
const columnsAccessor = [
...this.__columnsAccessor,
...dataframe.__columnsAccessor
...srcOffsets.map(i => dataframe.__columnsAccessor[i])
];
return new this.constructor(
dims,
columns,

View File

@@ -1,5 +1,8 @@
/*
Helpers for schema management
TODO: all this would be much more natural if done with a framework
like immutable.js
*/
import _ from "lodash";
@@ -31,8 +34,8 @@ export function indexEntireSchema(schema) {
return schema;
}
function _copy(schema) {
/* redux copy conventions - WARNING, only for modifyign obs annotations */
function _copyObsAnno(schema) {
/* redux copy conventions - WARNING, only for modifying obs annotations */
return {
...schema,
annotations: {
@@ -42,7 +45,17 @@ function _copy(schema) {
};
}
function _reindex(schema) {
function _copyObsLayout(schema) {
return {
...schema,
layout: {
...schema.layout,
obs: _.cloneDeep(schema.layout.obs)
}
};
}
function _reindexObsAnno(schema) {
/* reindex obs annotations ONLY */
schema.annotations.obsByName = fromEntries(
schema.annotations.obs.columns.map(v => [v.name, v])
@@ -50,18 +63,25 @@ function _reindex(schema) {
return schema;
}
function _reindexObsLayout(schema) {
schema.layout.obsByName = fromEntries(
schema.layout.obs.map(v => [v.name, v])
);
return schema;
}
export function removeObsAnnoColumn(schema, name) {
const newSchema = _copy(schema);
const newSchema = _copyObsAnno(schema);
newSchema.annotations.obs.columns = schema.annotations.obs.columns.filter(
v => v.name !== name
);
return _reindex(newSchema);
return _reindexObsAnno(newSchema);
}
export function addObsAnnoColumn(schema, name, defn) {
const newSchema = _copy(schema);
const newSchema = _copyObsAnno(schema);
newSchema.annotations.obs.columns.push(defn);
return _reindex(newSchema);
return _reindexObsAnno(newSchema);
}
export function removeObsAnnoCategory(schema, name, category) {
@@ -73,7 +93,7 @@ export function removeObsAnnoCategory(schema, name, category) {
const idx = categories.indexOf(category);
if (idx === -1) throw new Error("category does not exist");
const newSchema = _reindex(_copy(schema));
const newSchema = _reindexObsAnno(_copyObsAnno(schema));
/* remove category. Do not need to resort as this can't change presentation order */
newSchema.annotations.obsByName[name].categories.splice(idx, 1);
@@ -89,7 +109,7 @@ export function addObsAnnoCategory(schema, name, category) {
const idx = categories.indexOf(category);
if (idx !== -1) throw new Error("category already exists");
const newSchema = _reindex(_copy(schema));
const newSchema = _reindexObsAnno(_copyObsAnno(schema));
/* add category, retaining presentation sort order */
const catAnno = newSchema.annotations.obsByName[name];
@@ -99,3 +119,17 @@ export function addObsAnnoCategory(schema, name, category) {
]);
return newSchema;
}
export function addObsLayout(schema, layout) {
/* add or replace a layout */
const newSchema = _copyObsLayout(schema);
newSchema.layout.obs.push(layout);
return _reindexObsLayout(newSchema);
}
export function removeObsLayout(schema, name) {
/* remove a layout */
const newSchema = _copyObsLayout(schema);
newSchema.layout.obs = schema.layout.obs.filter(v => v.name !== name);
return _reindexObsLayout(newSchema);
}

View File

@@ -163,6 +163,10 @@ class LayoutObsAPI(Resource):
def get(self, data_adaptor):
return common_rest.layout_obs_get(request, data_adaptor)
@rest_get_data_adaptor
def put(self, data_adaptor):
return common_rest.layout_obs_put(request, data_adaptor)
def get_api_resources(bp_api):
api = Api(bp_api)

View File

@@ -13,7 +13,7 @@ import click
from server.common.utils import custom_format_warning
from server.common.utils import find_available_port, is_port_available, sort_options
from server.common.errors import DatasetAccessError
from server.data_common.matrix_loader import MatrixDataLoader, MatrixDataCacheManager
from server.data_common.matrix_loader import MatrixDataLoader, MatrixDataCacheManager, MatrixDataType
from server.common.annotations import AnnotationsLocalFile
from server.common.app_config import AppConfig
@@ -103,6 +103,14 @@ def config_args(func):
metavar="<text>",
help="Embedding name, eg, 'umap'. Repeat option for multiple embeddings. Defaults to all.",
)
@click.option(
"--experimental-enable-reembedding",
is_flag=True,
default=False,
show_default=False,
hidden=True,
help="Enable experimental on-demand re-embedding using UMAP. WARNING: may be very slow.",
)
@functools.wraps(func)
def wrapper(*args, **kwargs):
return func(*args, **kwargs)
@@ -278,6 +286,7 @@ def launch(
disable_diffexp,
experimental_annotations_ontology,
experimental_annotations_ontology_obo,
experimental_enable_reembedding,
):
"""Launch the cellxgene data viewer.
This web app lets you explore single-cell expression data.
@@ -317,6 +326,14 @@ def launch(
except DatasetAccessError as e:
raise click.ClickException(str(e))
if experimental_enable_reembedding:
if matrix_data_loader.matrix_data_type() != MatrixDataType.H5AD:
raise click.ClickException("--experimental-enable-reembedding is only supported with H5AD files.")
if backed:
raise click.ClickException(
"--experimental-enable-reembedding is not supported when run in --backed mode."
)
file_size = matrix_data_loader.file_size()
if file_size > BIG_FILE_SIZE_THRESHOLD:
click.echo(f"[cellxgene] Loading data from {basename(datapath)}, this may take a while...")
@@ -402,6 +419,7 @@ def launch(
var_names=var_names,
anndata_backed=backed,
disable_diffexp=disable_diffexp,
enable_reembedding=experimental_enable_reembedding,
)
matrix_data_cache_manager = MatrixDataCacheManager()

View File

@@ -32,6 +32,7 @@ class AppConfig(object):
self.max_category_items = 100
self.diffexp_lfc_cutoff = 0.01
self.disable_diffexp = False
self.enable_reembedding = False
self.anndata_backed = False
# TODO these options may not apply to all datasets in the multi dataset.
@@ -56,6 +57,7 @@ class AppConfig(object):
"var_names",
"anndata_backed",
"disable_diffexp",
"enable_reembedding",
]
self.update(inputs, kw)
@@ -80,7 +82,7 @@ class AppConfig(object):
# we have camalCase, hyphen-text, and underscore_text
# features
features = [f.todict() for f in data_adaptor.get_features().values()]
features = [f.todict() for f in data_adaptor.get_features(annotation)]
# display_names
title = self.get_title(data_adaptor)
@@ -105,6 +107,7 @@ class AppConfig(object):
"diffexp_lfc_cutoff": self.diffexp_lfc_cutoff,
"backed": self.anndata_backed,
"disable-diffexp": self.disable_diffexp,
"enable-reembedding": self.enable_reembedding,
"annotations": False,
"annotations_file": None,
"annotations_output_dir": None,

View File

@@ -165,7 +165,7 @@ def diffexp_obs_post(request, data_adaptor):
try:
diffexp = data_adaptor.diffexp_topN(set1_filter, set2_filter, count)
return make_response(diffexp, HTTPStatus.OK, {"Content-Type": "application/json"})
except (ValueError, FilterError) as e:
except (ValueError, DisabledFeatureError, FilterError) as e:
return make_response(str(e), HTTPStatus.BAD_REQUEST)
except JSONEncodingValueError as e:
# JSON encoding failure, usually due to bad data
@@ -188,3 +188,33 @@ def layout_obs_get(request, data_adaptor):
return make_response(str(e), HTTPStatus.INTERNAL_SERVER_ERROR)
except ValueError as e:
return make_response(str(e), HTTPStatus.INTERNAL_SERVER_ERROR)
def layout_obs_put(request, data_adaptor):
preferred_mimetype = request.accept_mimetypes.best_match(["application/octet-stream"])
if preferred_mimetype != "application/octet-stream":
return make_response(f"Unsupported MIME type '{request.accept_mimetypes}'", HTTPStatus.NOT_ACCEPTABLE)
if not data_adaptor.config.enable_reembedding:
return make_response(f"Computed embedding not supported.", HTTPStatus.BAD_REQUEST)
args = request.get_json()
filter = args["filter"] if args else None
if not filter:
return make_response("Error: obs filter is required", HTTPStatus.BAD_REQUEST)
method = args["method"] if args else "umap"
try:
schema, fbs = data_adaptor.compute_embedding(method, filter)
return make_response(
fbs,
HTTPStatus.OK,
{
"Content-Type": "application/octet-stream",
"CxG-Schema": json.dumps(schema),
"Access-Control-Expose-Headers": "CxG-Schema",
},
)
except NotImplementedError as e:
return make_response(str(e), HTTPStatus.NOT_IMPLEMENTED)
except (ValueError, DisabledFeatureError, FilterError) as e:
return make_response(str(e), HTTPStatus.BAD_REQUEST)

49
server/compute/scanpy.py Normal file
View File

@@ -0,0 +1,49 @@
import importlib
"""
Wrapper for various scanpy modules. Will raise NotImplementedError if the scanpy
module is not installed/available
"""
def get_scanpy_module():
try:
sc = importlib.import_module("scanpy")
# Future: we could enforce versions here, eg, lookat sc.__version__
return sc
except ModuleNotFoundError:
raise NotImplementedError("Please install scanpy to enable UMAP re-embedding")
except Exception as e:
# will capture other ImportError corner cases
raise NotImplementedError(str(e))
def scanpy_umap(adata, obs_mask=None, pca_options={}, neighbors_options={}, umap_options={}):
"""
Given adata and an obs mask, return a new embedding for adata[obs_mask, :]
as an ndarray of shape (len(obs_mask), N), where N>=2.
Do NOT mutate adata.
"""
# backed mode is incompatible with the current implementation
if adata.isbacked:
raise NotImplementedError("Backed mode is incompatible with re-embedding")
# safely get scanpy module, which may not be present.
sc = get_scanpy_module()
# https://github.com/theislab/anndata/issues/311
obs_mask = slice(None) if obs_mask is None else obs_mask
adata = adata[obs_mask, :].copy()
for k in list(adata.obsm.keys()):
del adata.obsm[k]
for k in list(adata.uns.keys()):
del adata.uns[k]
sc.pp.pca(adata, zero_center=None, n_comps=min(adata.n_obs - 1, 50), **pca_options)
sc.pp.neighbors(adata, **neighbors_options)
sc.tl.umap(adata, **umap_options)
return adata.obsm["X_umap"]

View File

@@ -1,17 +1,21 @@
import warnings
import numpy as np
import pandas as pd
from pandas.core.dtypes.dtypes import CategoricalDtype
import anndata
from scipy import sparse
from packaging import version
from datetime import datetime
from server_timing import Timing as ServerTiming
from server.data_common.data_adaptor import DataAdaptor
from server.data_common.fbs.matrix import encode_matrix_fbs
from server.common.utils import series_to_schema
from server.common.constants import Axis, MAX_LAYOUTS
from server.common.errors import PrepareError, DatasetAccessError
from server.common.errors import PrepareError, DatasetAccessError, FilterError
from server.common.data_locator import DataLocator
from server.compute.scanpy import scanpy_umap
anndata_version = version.parse(str(anndata.__version__)).release
@@ -261,7 +265,10 @@ class AnndataAdaptor(DataAdaptor):
return encode_matrix_fbs(df, col_idx=df.columns)
def get_embedding_names(self):
""" function:
"""
Return pre-computed embeddings.
function:
a) generate list of default layouts
b) validate layouts are legal. remove/warn on any that are not
c) cap total list of layouts at global const MAX_LAYOUTS
@@ -294,6 +301,30 @@ class AnndataAdaptor(DataAdaptor):
full_embedding = self.data.obsm[f"X_{ename}"]
return full_embedding[:, 0:dims]
def compute_embedding(self, method, obsFilter):
if Axis.VAR in obsFilter:
raise FilterError("Observation filters may not contain variable conditions")
if method != "umap":
raise NotImplementedError(f"re-embedding method {method} is not available.")
try:
shape = self.get_shape()
obs_mask = self._axis_filter_to_mask(Axis.OBS, obsFilter["obs"], shape[0])
except (KeyError, IndexError) as e:
raise FilterError(f"Error parsing filter: {e}") from e
with ServerTiming.time("layout.compute"):
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
# embedding name generated by this backed.
name = f"reembed:{method}_{datetime.now().isoformat(timespec='milliseconds')}"
dims = [f"{name}_0", f"{name}_1"]
df = pd.DataFrame(normalized_layout, columns=dims)
fbs = encode_matrix_fbs(df, col_idx=df.columns, row_idx=None)
schema = {"name": name, "type": "float32", "dims": dims}
return (schema, fbs)
def get_X_array(self, obs_mask=None, var_mask=None):
if obs_mask is None:
obs_mask = slice(None)

View File

@@ -57,12 +57,18 @@ class DataAdaptor(metaclass=ABCMeta):
@abstractmethod
def get_embedding_names(self):
"""return a list of embedding names"""
"""return a list of pre-computed embedding names"""
pass
@abstractmethod
def get_embedding_array(self, ename, dims=2):
"""return an numpy array for the given embedding name."""
"""return an numpy array for the given pre-computed embedding name."""
pass
@abstractmethod
def compute_embedding(self, method, filter):
"""compute a new embedding on the specified obs subset, and return a
tuple of (schema, fbs)."""
pass
@abstractmethod
@@ -126,21 +132,15 @@ class DataAdaptor(metaclass=ABCMeta):
"""
pass
def get_features(self):
features = {}
features["cluster"] = AppFeature("/cluster/")
if self.get_embedding_names():
# TODO handle "var" when gene layout becomes available
features["layout_obs"] = AppFeature("/layout/obs", available=True)
else:
features["layout_obs"] = AppFeature("/layout/obs")
if self.config.disable_diffexp:
features["diffexp"] = AppFeature("/diffexp/")
else:
features["diffexp"] = AppFeature("/diffexp/", available=True)
def get_features(self, annotations=None):
"""Return list of features, to return as part of the config route"""
features = [
AppFeature("/cluster/", method="POST", available=False),
AppFeature("/layout/obs", method="GET", available=self.get_embedding_names() is not None),
AppFeature("/layout/obs", method="PUT", available=self.config.enable_reembedding),
AppFeature("/diffexp/", method="POST", available=not self.config.disable_diffexp),
AppFeature("/annotations/obs", method="PUT", available=annotations is not None),
]
return features
def update_parameters(self, parameters):
@@ -294,6 +294,25 @@ class DataAdaptor(metaclass=ABCMeta):
except ValueError:
raise JSONEncodingValueError("Error encoding differential expression to JSON")
@staticmethod
def normalize_embedding(embedding):
"""Normalize embedding layout to meet client assumptions.
Embedding is an ndarray, shape (n_obs, n)., where n is normally 2
"""
# scale isotropically
min = embedding.min(axis=0)
max = embedding.max(axis=0)
scale = np.amax(max - min)
normalized_layout = (embedding - min) / scale
# translate to center on both axis
translate = 0.5 - ((max - min) / scale / 2)
normalized_layout = normalized_layout + translate
normalized_layout = normalized_layout.astype(dtype=np.float32)
return normalized_layout
def layout_to_fbs_matrix(self):
""" same as layout, except returns a flatbuffer """
"""
@@ -312,18 +331,7 @@ class DataAdaptor(metaclass=ABCMeta):
with ServerTiming.time(f"layout.query"):
for ename in embeddings:
embedding = self.get_embedding_array(ename, 2)
# scale isotropically
min = embedding.min(axis=0)
max = embedding.max(axis=0)
scale = np.amax(max - min)
normalized_layout = (embedding - min) / scale
# translate to center on both axis
translate = 0.5 - ((max - min) / scale / 2)
normalized_layout = normalized_layout + translate
normalized_layout = normalized_layout.astype(dtype=np.float32)
normalized_layout = DataAdaptor.normalize_embedding(embedding)
layout_data.append(pd.DataFrame(normalized_layout, columns=[f"{ename}_0", f"{ename}_1"]))
with ServerTiming.time(f"layout.encode"):

View File

@@ -164,6 +164,9 @@ class CxgAdaptor(DataAdaptor):
array = self.open_array(f"emb/{ename}")
return array[:, 0:dims]
def compute_embedding(self, method, filter):
raise NotImplementedError("CXG does not yet support re-embedding")
def get_X_array(self, obs_mask=None, var_mask=None):
obs_items = self._convert_mask(obs_mask)
var_items = self._convert_mask(var_mask)

View File

@@ -3,6 +3,7 @@ from os import path
import pytest
import time
import unittest
import sys
import server.test.decode_fbs as decode_fbs
from parameterized import parameterized_class
@@ -97,7 +98,21 @@ class AdaptorTest(unittest.TestCase):
self.data._create_schema()
def test_config(self):
self.assertEqual(self.data.get_features()["layout_obs"].available, True)
features = self.data.get_features(annotations=None)
# test each for singular presence and accuracy of available flag
def check_feature(method, path, available):
feature = list(
filter(lambda f: f.method == method and f.path == path and f.available == available, features)
)
self.assertIsNotNone(feature)
self.assertEqual(len(feature), 1)
check_feature("POST", "/cluster/", False)
check_feature("POST", "/diffexp/", not self.data.config.disable_diffexp)
check_feature("GET", "/layout/obs", True)
check_feature("PUT", "/layout/obs", self.data.config.enable_reembedding)
check_feature("PUT", "/annotations/obs", False)
def test_layout(self):
fbs = self.data.layout_to_fbs_matrix()
@@ -185,3 +200,39 @@ class AdaptorTest(unittest.TestCase):
self.assertEqual(data["n_rows"], 2638)
self.assertEqual(data["n_cols"], 3)
self.assertTrue((data["col_idx"] == [15, 1818, 1837]).all())
def test_compute_embedding(self):
filter = {"obs": {"index": [[0, 100]]}}
# Verify that we correctly handle the case where we lack scanpy
import unittest.mock
with unittest.mock.patch.dict(sys.modules, {"scanpy": None}):
with self.assertRaises(NotImplementedError):
self.data.compute_embedding("umap", filter)
# if we happen to have scanpy, test the full API, else punt
import importlib
scanpy_spec = importlib.util.find_spec("scanpy")
if scanpy_spec is None:
print("Skipping compute_embedding test as ScanPy not installed")
return
# this feature is unsupported in backed mode, and we expect an error
if self.data.data.isbacked:
with self.assertRaises(NotImplementedError):
self.data.compute_embedding("umap", filter)
return
(schema, fbs) = self.data.compute_embedding("umap", filter)
self.assertIsInstance(schema["name"], str)
name = schema["name"]
self.assertEqual(schema["type"], "float32")
self.assertEqual(schema["dims"], [f"{name}_0", f"{name}_1"])
emb = decode_fbs.decode_matrix_FBS(fbs)
self.assertEqual(emb["n_rows"], 100)
self.assertEqual(emb["n_cols"], 2)
self.assertEqual(emb["col_idx"], [f"{name}_0", f"{name}_1"])

View File

@@ -44,7 +44,7 @@ class EndPoints(object):
result_data = result.json()
self.assertIn("library_versions", result_data["config"])
self.assertEqual(result_data["config"]["displayNames"]["dataset"], "pbmc3k")
self.assertEqual(len(result_data["config"]["features"]), 3)
self.assertEqual(len(result_data["config"]["features"]), 5)
def test_get_layout_fbs(self):
endpoint = "layout/obs"

View File

@@ -136,3 +136,20 @@ class WritableAnnotationTest(unittest.TestCase):
all_col_schema["cat_B"],
{"name": "cat_B", "type": "categorical", "categories": ["label_B"], "writable": True},
)
def test_config(self):
features = self.data.get_features(self.annotations)
# test each for singular presence and accuracy of available flag
def check_feature(method, path, available):
feature = list(
filter(lambda f: f.method == method and f.path == path and f.available == available, features)
)
self.assertIsNotNone(feature)
self.assertEqual(len(feature), 1)
check_feature("POST", "/cluster/", False)
check_feature("POST", "/diffexp/", not self.data.config.disable_diffexp)
check_feature("GET", "/layout/obs", True)
check_feature("PUT", "/layout/obs", self.data.config.enable_reembedding)
check_feature("PUT", "/annotations/obs", True)