From efe3bf7a72aceee65887bd53b440d17a56129875 Mon Sep 17 00:00:00 2001 From: Emanuele Bezzi Date: Tue, 7 Dec 2021 19:42:40 -0500 Subject: [PATCH] Parametrization --- client/src/actions/index.js | 2 + client/src/actions/spatial.js | 35 +++++++++++ client/src/components/app.js | 2 + .../components/graph/drawSpatialImageRegl.js | 26 +------- client/src/components/graph/graph.js | 60 +++++++++++++++---- client/src/reducers/index.js | 3 + client/src/reducers/spatial.js | 34 +++++++++++ server/app/app.py | 6 ++ server/data_anndata/anndata_adaptor.py | 34 +++++++++++ server/data_common/data_adaptor.py | 11 +++- 10 files changed, 177 insertions(+), 36 deletions(-) create mode 100644 client/src/actions/spatial.js create mode 100644 client/src/reducers/spatial.js diff --git a/client/src/actions/index.js b/client/src/actions/index.js index 5c5f495d..7f3dc64b 100644 --- a/client/src/actions/index.js +++ b/client/src/actions/index.js @@ -8,6 +8,7 @@ import { import { loadUserColorConfig } from "../util/stateManager/colorHelpers"; import * as selnActions from "./selection"; import * as annoActions from "./annotation"; +import * as spatialActions from "./spatial"; import * as viewActions from "./viewStack"; import * as embActions from "./embedding"; import * as genesetActions from "./geneset"; @@ -272,4 +273,5 @@ export default { genesetDelete: genesetActions.genesetDelete, genesetAddGenes: genesetActions.genesetAddGenes, genesetDeleteGenes: genesetActions.genesetDeleteGenes, + requestSpatialMetadata: spatialActions.requestSpatialMetadata, }; diff --git a/client/src/actions/spatial.js b/client/src/actions/spatial.js new file mode 100644 index 00000000..7a928568 --- /dev/null +++ b/client/src/actions/spatial.js @@ -0,0 +1,35 @@ +import * as globals from "../globals"; + +export const requestSpatialMetadata = () => async (dispatch) => { + dispatch({ type: "request spatial metadata started" }); + try { + const res = await fetch( + `${globals.API.prefix}${globals.API.version}spatial/meta`, + { + method: "GET", + headers: new Headers({ + Accept: "application/json", + "Content-Type": "application/json", + }), + credentials: "include", + } + ); + + if (!res.ok || res.headers.get("Content-Type") !== "application/json") { + return null; // TODO need a dispatch //dispatchDiffExpErrors(dispatch, res); + } + + const response = await res.json(); + + /* then send the success case action through */ + return dispatch({ + type: "request spatial metadata success", + data: response, + }); + } catch (error) { + return dispatch({ + type: "request spatial metadata error", + error, + }); + } +}; diff --git a/client/src/components/app.js b/client/src/components/app.js index d3b78761..f968db12 100644 --- a/client/src/components/app.js +++ b/client/src/components/app.js @@ -23,6 +23,8 @@ class App extends React.Component { componentDidMount() { const { dispatch } = this.props; + dispatch(actions.requestSpatialMetadata()); + /* listen for url changes, fire one when we start the app up */ window.addEventListener("popstate", this._onURLChanged); this._onURLChanged(); diff --git a/client/src/components/graph/drawSpatialImageRegl.js b/client/src/components/graph/drawSpatialImageRegl.js index 3046184a..37c15faf 100644 --- a/client/src/components/graph/drawSpatialImageRegl.js +++ b/client/src/components/graph/drawSpatialImageRegl.js @@ -45,35 +45,15 @@ export default function drawSpatialImageRegl(regl) { attributes: { a_texCoord: [0.0, 0.0, 1.0, 0.0, 0.0, 1.0, 0.0, 1.0, 1.0, 0.0, 1.0, 1.0], - // a_position: [ - // 10, 0, - // 10 + regl.prop("img_width"), 0, - // 10, 0 + regl.prop("img_height"), - // 10, 0 + regl.prop("img_height"), - // 10 + regl.prop("img_width"), 0, - // 10 + regl.prop("img_width"), 0 + regl.prop("img_height"), - // ], - a_position: [ - 0, - 0, - 0 + 1921, - 0, - 0, - 0 + 2000, - 0, - 0 + 2000, - 0 + 1921, - 0, - 0 + 1921, - 0 + 2000, - ], + a_position: regl.prop("rectCoords"), }, uniforms: { projView: regl.prop("projView"), u_image: regl.prop("spatialImageAsTexture"), color: [1, 0, 0, 1], - u_resolution: [1921, 2000], + u_resolution: [regl.prop("imageWidth"), regl.prop("imageHeight")], + image_width: regl.prop("imageWidth"), // translate: }, diff --git a/client/src/components/graph/graph.js b/client/src/components/graph/graph.js index 43923c04..59a826a4 100644 --- a/client/src/components/graph/graph.js +++ b/client/src/components/graph/graph.js @@ -78,6 +78,7 @@ function createModelTF() { colors: state.colors, pointDilation: state.pointDilation, genesets: state.genesets.genesets, + spatial: state.spatial.metadata, })) class Graph extends React.Component { static createReglState(canvas) { @@ -110,19 +111,20 @@ class Graph extends React.Component { return !shallowEqual(props.watchProps, prevProps.watchProps); } - computePointPositions = memoize((X, Y, modelTF) => { + computePointPositions = memoize((X, Y, modelTF, spatialMetadata) => { /* compute the model coordinate for each point */ console.log({ X }, { Y }); const positions = new Float32Array(2 * X.length); for (let i = 0, len = X.length; i < len; i += 1) { - const p = vec2.fromValues(X[i], Y[i]); + // TODO: Introduce the feature flag here + // const p = vec2.fromValues(X[i], Y[i]); + const p = this.rescalePointForSpatial(X[i], Y[i], spatialMetadata); vec2.transformMat3(p, p, modelTF); positions[2 * i] = p[0]; positions[2 * i + 1] = p[1]; } - console.log({ transformed: positions }); return positions; }); @@ -238,6 +240,7 @@ class Graph extends React.Component { colorBuffer: null, flagBuffer: null, drawSpatialImage: null, + spatial: null, // component rendering derived state - these must stay synchronized // with the reducer state they were generated from. @@ -452,6 +455,33 @@ class Graph extends React.Component { }); } + rescalePointForSpatial = (x, y, spatialMetadata) => { + // console.log({spatialMetadata}); + const translate = vec2.fromValues( + spatialMetadata.inverseTranslate[0], + spatialMetadata.inverseTranslate[1] + ); + const min = vec2.fromValues( + spatialMetadata.inverseMin[0], + spatialMetadata.inverseMin[1] + ); + const scalefactor = spatialMetadata.scaleref; + const wh = vec2.fromValues( + spatialMetadata.imageWidth, + spatialMetadata.imageHeight + ); + + // Apply the inverse transform + const p = vec2.fromValues(x, y); + vec2.sub(p, p, translate); + vec2.scale(p, p, spatialMetadata.inverseScale); + vec2.add(p, p, min); + + vec2.scale(p, p, scalefactor); + vec2.div(p, p, wh); + return p; + }; + setReglCanvas = (canvas) => { this.reglCanvas = canvas; this.setState({ @@ -534,9 +564,12 @@ class Graph extends React.Component { crossfilter, pointDilation, viewport, + spatial, } = props.watchProps; const { modelTF } = this.state; + console.log({ spatial }); + const [layoutDf, colorDf, pointDilationDf] = await this.fetchData( annoMatrix, layoutChoice, @@ -547,7 +580,7 @@ class Graph extends React.Component { const { currentDimNames } = layoutChoice; const X = layoutDf.col(currentDimNames[0]).asArray(); const Y = layoutDf.col(currentDimNames[1]).asArray(); - const positions = this.computePointPositions(X, Y, modelTF); + const positions = this.computePointPositions(X, Y, modelTF, spatial.data); const colorTable = this.updateColorTable(colorsProp, colorDf); const colors = this.computePointColors(colorTable.rgb); @@ -579,6 +612,7 @@ class Graph extends React.Component { flags, width, height, + spatial, }; }; @@ -760,7 +794,6 @@ class Graph extends React.Component { const { positions, colors, flags, height, width } = asyncProps; this.cachedAsyncProps = asyncProps; const { pointBuffer, colorBuffer, flagBuffer } = this.state; - console.log({ pos2: positions }); let needToRenderCanvas = false; if (height !== prevAsyncProps?.height || width !== prevAsyncProps?.width) { @@ -824,20 +857,21 @@ class Graph extends React.Component { projectionTF, drawSpatialImage ) { - const { annoMatrix } = this.props; + const { annoMatrix, spatial } = this.props; if (!this.reglCanvas || !annoMatrix) return; const { schema } = annoMatrix; const cameraTF = camera.view(); const projView = mat3.multiply(mat3.create(), projectionTF, cameraTF); const { width, height } = this.reglCanvas; + const imW = spatial.data.imageWidth; + const imH = spatial.data.imageHeight; + regl.poll(); regl.clear({ depth: 1, color: [0, 0, 0, 0], }); - console.log({ pointBuffer }); - console.log({ projView }); drawPoints({ distance: camera.distance(), color: colorBuffer, @@ -850,8 +884,9 @@ class Graph extends React.Component { }); drawSpatialImage({ projView, - img_width: this.spatialImage.width, - img_height: this.spatialImage.height, + imageWidth: imW, + imageHeight: imH, + rectCoords: [0, 0, imW, 0, 0, imH, 0, imH, imW, 0, imW, imH], spatialImageAsTexture: regl.texture({ data: this.spatialImage, wrapS: "clamp", @@ -869,10 +904,14 @@ class Graph extends React.Component { layoutChoice, pointDilation, crossfilter, + spatial, } = this.props; const { modelTF, projectionTF, camera, viewport, regl } = this.state; const cameraTF = camera?.view()?.slice(); + console.log("---RENDER"); + console.log({ props: this.props }); + return (
diff --git a/client/src/reducers/index.js b/client/src/reducers/index.js index 29ace171..05da1b88 100644 --- a/client/src/reducers/index.js +++ b/client/src/reducers/index.js @@ -11,6 +11,7 @@ import continuousSelection from "./continuousSelection"; import graphSelection from "./graphSelection"; import colors from "./colors"; import differential from "./differential"; +import spatial from "./spatial"; import layoutChoice from "./layoutChoice"; import controls from "./controls"; import annotations from "./annotations"; @@ -38,6 +39,7 @@ const Reducer = undoable( ["colors", colors], ["controls", controls], ["differential", differential], + ["spatial", spatial], ["centroidLabels", centroidLabels], ["pointDilation", pointDialation], ["autosave", autosave], @@ -51,6 +53,7 @@ const Reducer = undoable( "colors", "controls", "differential", + "spatial", "layoutChoice", "centroidLabels", "genesets", diff --git a/client/src/reducers/spatial.js b/client/src/reducers/spatial.js new file mode 100644 index 00000000..31007594 --- /dev/null +++ b/client/src/reducers/spatial.js @@ -0,0 +1,34 @@ +const Spatial = ( + state = { + loading: null, + error: null, + metadata: null, + }, + action +) => { + switch (action.type) { + case "request spatial metadata started": + return { + ...state, + loading: true, + error: null, + }; + case "request spatial metadata success": + return { + ...state, + error: null, + loading: false, + metadata: action, + }; + case "request spatial metadata error": + return { + ...state, + loading: false, + error: action.data, + }; + default: + return state; + } +}; + +export default Spatial; diff --git a/server/app/app.py b/server/app/app.py index 2e3e418a..9c270bbe 100644 --- a/server/app/app.py +++ b/server/app/app.py @@ -195,6 +195,11 @@ class SpatialImageAPI(Resource): def get(self, data_adaptor): return common_rest.spatial_image_get(request, data_adaptor) +class SpatialMetaAPI(Resource): + @rest_get_data_adaptor + def get(self, data_adaptor): + return data_adaptor.get_spatial_metadata() + def get_api_base_resources(bp_base): """Add resources that are accessed from the api url""" @@ -229,6 +234,7 @@ def get_api_dataroot_resources(bp_dataroot): add_resource(LayoutObsAPI, "/layout/obs") # Spatial routes add_resource(SpatialImageAPI, "/spatial/image") + add_resource(SpatialMetaAPI, "/spatial/meta") return api diff --git a/server/data_anndata/anndata_adaptor.py b/server/data_anndata/anndata_adaptor.py index 46f66412..5b33cfb1 100644 --- a/server/data_anndata/anndata_adaptor.py +++ b/server/data_anndata/anndata_adaptor.py @@ -277,6 +277,40 @@ class AnndataAdaptor(DataAdaptor): def get_spatial(self): return self.data.uns["spatial"] + def get_spatial_metadata(self): + spatial = self.get_spatial() + + resolution = "hires" + + if len(list(spatial)) == 0: + raise Exception("uns does not have spatial information") + + library_id = list(spatial)[0] + + if "images" not in spatial[library_id]: + raise Exception("spatial information does not contain images") + + if resolution not in spatial[library_id]["images"]: + raise Exception(f"spatial information does not contain requested resolution '{resolution}'") + + scaleref = spatial[library_id]["scalefactors"][f"tissue_{resolution}_scalef"] + (h, w, _) = spatial[library_id]["images"][resolution].shape + + A = self.data.obsm["X_spatial"] + min = np.nanmin(A, axis=0) + max = np.nanmax(A, axis=0) + scale = np.amax(max - min) + translate = 0.5 - ((max - min) / scale / 2) + + return { + "imageWidth": w, + "imageHeight": h, + "scaleref": scaleref, + "inverseScale": int(scale), + "inverseTranslate": translate.tolist(), + "inverseMin": min.tolist(), + } + def get_embedding_names(self): """ Return pre-computed embeddings. diff --git a/server/data_common/data_adaptor.py b/server/data_common/data_adaptor.py index 90e07dea..fa3113dd 100644 --- a/server/data_common/data_adaptor.py +++ b/server/data_common/data_adaptor.py @@ -363,10 +363,15 @@ class DataAdaptor(metaclass=ABCMeta): print(f"scale {scale}, translate {translate}") - A = embedding * 0.17011142 - A = np.column_stack([A[:, 0] / 1921, A[:, 1] / 2000]) + # if True: # if visium + # self.data.uns["spatial"] - normalized_layout = A.astype(dtype=np.float32) + + # adata.uns["spatial"]['V1_Adult_Mouse_Brain']["scalefactors"]["tissue_hires_scalef"] + # A = embedding * 0.17011142 + # A = np.column_stack([A[:, 0] / 1921, A[:, 1] / 2000]) + + # normalized_layout = A.astype(dtype=np.float32) return normalized_layout def layout_to_fbs_matrix(self, fields):