Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
82 changes: 82 additions & 0 deletions bats_ai/core/views/nabat/nabat_recording.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import base64
import json
import logging
from typing import TYPE_CHECKING, Any

from django.conf import settings
from django.db import transaction
Expand All @@ -17,12 +18,17 @@
from bats_ai.core.models import ProcessingTask, ProcessingTaskType, Species
from bats_ai.core.models.nabat import (
NABatCompressedSpectrogram,
NABatPulseMetadata,
NABatRecording,
NABatRecordingAnnotation,
)
from bats_ai.core.tasks.nabat.nabat_data_retrieval import nabat_recording_initialize
from bats_ai.core.views.species import SpeciesSchema

if TYPE_CHECKING:
from bats_ai.core.views.recording import PulseMetadataSlopesSchema


logger = logging.getLogger(__name__)
router = RouterPaginated()

Expand Down Expand Up @@ -590,3 +596,79 @@ def delete_recording_annotation(
# Check permission
annotation.delete()
return "Recording annotation deleted successfully."


class NABatPulseContourSchema(Schema):
id: int | None
index: int
bounding_box: Any
contours: list

@classmethod
def from_orm(cls, obj: NABatPulseMetadata):
return cls(
id=obj.id,
index=obj.index,
contours=obj.contours if obj.contours is not None else [],
bounding_box=json.loads(obj.bounding_box.geojson),
)


class NABatPulseMetadataSchema(Schema):
id: int | None
index: int
curve: list[list[float]] | None = None
char_freq: list[float] | None = None
knee: list[float] | None = None
heel: list[float] | None = None
slopes: PulseMetadataSlopesSchema | None = None

@classmethod
def from_orm(cls, obj: NABatPulseMetadata):
def point_to_list(pt):
if pt is None:
return None
return [pt.x, pt.y]

def linestring_to_list(ls):
if ls is None:
return None
return [[c[0], c[1]] for c in ls.coords]

return cls(
id=obj.id,
index=obj.index,
curve=linestring_to_list(obj.curve),
char_freq=point_to_list(obj.char_freq),
knee=point_to_list(obj.knee),
heel=point_to_list(obj.heel),
slopes=obj.slopes,
)


@router.get("/{pk}/pulse_contours", auth=None)
def get_pulse_contours(request: HttpRequest, pk: int, api_token: str):
recording = get_object_or_404(NABatRecording, pk=pk)

email_or_response = get_email_if_authorized(request, api_token, recording.recording_id)
if isinstance(email_or_response, JsonResponse):
return email_or_response

computed_pulse_annotation_qs = NABatPulseMetadata.objects.filter(
nabat_recording=recording
).order_by("index")
return [NABatPulseContourSchema.from_orm(pulse) for pulse in computed_pulse_annotation_qs]


@router.get("/{pk}/pulse_metadata", auth=None)
def get_pulse_data(request: HttpRequest, pk: int, api_token: str):
recording = get_object_or_404(NABatRecording, pk=pk)

email_or_response = get_email_if_authorized(request, api_token, recording.recording_id)
if isinstance(email_or_response, JsonResponse):
return email_or_response

computed_pulse_annotation_qs = NABatPulseMetadata.objects.filter(
nabat_recording=recording
).order_by("index")
return [NABatPulseMetadataSchema.from_orm(pulse) for pulse in computed_pulse_annotation_qs]
20 changes: 20 additions & 0 deletions client/src/api/NABatApi.ts
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,8 @@ import {
type Spectrogram,
type UpdateFileAnnotation,
type Species,
type ComputedPulseContour,
type PulseMetadata,
} from "./api";

export interface NABatRecordingCompleteResponse {
Expand Down Expand Up @@ -275,6 +277,22 @@ async function exportNABatAnnotations(
return response.data;
}

async function getNabatPulseContours(recordingId: string, apiToken: string) {
const result = await axiosInstance.get<ComputedPulseContour[]>(
`nabat/recording/${recordingId}/pulse_contours`,
{ params: { api_token: apiToken } },
);
return result.data;
}

async function getNabatPulseMetadata(recordingId: string, apiToken: string) {
const result = await axiosInstance.get<PulseMetadata[]>(
`nabat/recording/${recordingId}/pulse_metadata`,
{ params: { api_token: apiToken } },
);
return result.data;
}

export {
postNABatRecording,
getNABatSpectrogram,
Expand All @@ -292,4 +310,6 @@ export {
getNABatConfigurationRecordings,
exportNABatAnnotations,
adminNaBatUpdateSpecies,
getNabatPulseContours,
getNabatPulseMetadata,
};
12 changes: 11 additions & 1 deletion client/src/components/PulseMetadataButton.vue
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ import { defineComponent, ref } from "vue";
import usePulseMetadata, {
PULSE_METADATA_LABELS_OPTIONS,
} from "@use/usePulseMetadata";
import useState from "@use/useState";

export default defineComponent({
name: "PulseMetadataButton",
Expand All @@ -22,6 +23,7 @@ export default defineComponent({
viewPulseMetadataLayer,
toggleViewPulseMetadataLayer,
loadPulseMetadata,
loadNabatPulseMetadata,
pulseMetadataList,
pulseMetadataLoading,
pulseMetadataLineColor,
Expand All @@ -35,10 +37,18 @@ export default defineComponent({
pulseMetadataLabels,
pulseMetadataDurationFreqLineColor,
} = usePulseMetadata();
const { isNaBat, nabatApiToken } = useState();

const togglePulseMetadata = async () => {
if (pulseMetadataList.value.length === 0 && props.recordingId != null) {
await loadPulseMetadata(Number(props.recordingId));
if (isNaBat()) {
await loadNabatPulseMetadata(
String(props.recordingId),
nabatApiToken.value,
);
} else {
await loadPulseMetadata(Number(props.recordingId));
}
}
toggleViewPulseMetadataLayer();
};
Expand Down
36 changes: 24 additions & 12 deletions client/src/components/SpectrogramViewer.vue
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,10 @@ export default defineComponent({
type: Array as PropType<HTMLImageElement[]>,
default: () => [],
},
maskLoaded: {
type: Boolean,
default: false,
},
waveplotImages: {
type: Array as PropType<HTMLImageElement[]>,
default: () => [],
Expand Down Expand Up @@ -496,18 +500,26 @@ export default defineComponent({
}
});

watch([viewMaskOverlay, maskOverlayOpacity, () => props.maskImages], () => {
if (viewMaskOverlay.value && props.maskImages.length) {
geoJS.drawMaskImages(
props.maskImages,
scaledWidth.value,
scaledHeight.value,
maskOverlayOpacity.value,
);
} else {
geoJS.clearMaskQuadFeatures(true);
}
});
watch(
[
viewMaskOverlay,
maskOverlayOpacity,
() => props.maskImages,
() => props.maskLoaded,
],
() => {
if (viewMaskOverlay.value && props.maskImages.length) {
geoJS.drawMaskImages(
props.maskImages,
scaledWidth.value,
scaledHeight.value,
maskOverlayOpacity.value,
);
} else {
geoJS.clearMaskQuadFeatures(true);
}
},
);

watch([showWaveplot], () => {
resetViewerBounds(false);
Expand Down
42 changes: 27 additions & 15 deletions client/src/components/ThumbnailViewer.vue
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,10 @@ export default defineComponent({
type: Array as PropType<HTMLImageElement[]>,
default: () => [],
},
maskLoaded: {
type: Boolean,
default: false,
},
waveplotImages: {
type: Array as PropType<HTMLImageElement[]>,
default: () => [],
Expand Down Expand Up @@ -259,21 +263,29 @@ export default defineComponent({
drawWaveplotIfEnabled(finalWidth, finalHeight);
});

watch([viewMaskOverlay, maskOverlayOpacity, () => props.maskImages], () => {
const { width, height } = getImageDimensions(props.images);
const finalWidth = scaledWidth.value || width;
const finalHeight = scaledHeight.value || height;
if (viewMaskOverlay.value && props.maskImages.length) {
geoJS.drawMaskImages(
props.maskImages,
finalWidth,
finalHeight,
maskOverlayOpacity.value,
);
} else {
geoJS.clearMaskQuadFeatures(true);
}
});
watch(
[
viewMaskOverlay,
maskOverlayOpacity,
() => props.maskImages,
() => props.maskLoaded,
],
() => {
const { width, height } = getImageDimensions(props.images);
const finalWidth = scaledWidth.value || width;
const finalHeight = scaledHeight.value || height;
if (viewMaskOverlay.value && props.maskImages.length) {
geoJS.drawMaskImages(
props.maskImages,
finalWidth,
finalHeight,
maskOverlayOpacity.value,
);
} else {
geoJS.clearMaskQuadFeatures(true);
}
},
);

watch(viewWaveplot, () => {
const { width, height } = getImageDimensions(props.images);
Expand Down
19 changes: 17 additions & 2 deletions client/src/components/geoJS/LayerManager.vue
Original file line number Diff line number Diff line change
Expand Up @@ -101,13 +101,17 @@ export default defineComponent({
contoursEnabled,
contourOpacity,
loadContours,
loadNabatContours,
isNaBat,
nabatApiToken,
computedPulseContours,
transparencyThreshold,
} = useState();
const {
viewPulseMetadataLayer,
pulseMetadataList,
loadPulseMetadata,
loadNabatPulseMetadata,
clearPulseMetadata,
pulseMetadataLineColor,
pulseMetadataLineSize,
Expand Down Expand Up @@ -595,7 +599,11 @@ export default defineComponent({
return;
}
if (computedPulseContours.value.length === 0) {
await loadContours(new Number(props.recordingId) as number);
if (isNaBat()) {
await loadNabatContours(props.recordingId);
} else {
await loadContours(new Number(props.recordingId) as number);
}
}
if (!contourLayer) {
contourLayer = new ContourLayer(
Expand Down Expand Up @@ -636,7 +644,14 @@ export default defineComponent({
if (!props.recordingId || !props.spectroInfo?.compressedWidth) return;
if (viewPulseMetadataLayer.value) {
if (pulseMetadataList.value.length === 0) {
await loadPulseMetadata(Number(props.recordingId));
if (isNaBat()) {
await loadNabatPulseMetadata(
props.recordingId,
nabatApiToken.value,
);
} else {
await loadPulseMetadata(Number(props.recordingId));
}
}
if (!pulseMetadataLayer) {
pulseMetadataLayer = new PulseMetadataLayer(
Expand Down
14 changes: 14 additions & 0 deletions client/src/use/usePulseMetadata.ts
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import { ref, type Ref, watch } from "vue";
import { getPulseMetadata, type PulseMetadata } from "../api/api";
import { getNabatPulseMetadata } from "@/api/NABatApi";

const STORAGE_KEY = "pulseMetadata";

Expand Down Expand Up @@ -79,6 +80,18 @@ async function loadPulseMetadata(recordingId: number) {
}
}

async function loadNabatPulseMetadata(recordingId: string, apiToken: string) {
pulseMetadataLoading.value = true;
try {
pulseMetadataList.value = await getNabatPulseMetadata(
recordingId,
apiToken,
);
} finally {
pulseMetadataLoading.value = false;
}
}

function clearPulseMetadata() {
pulseMetadataList.value = [];
}
Expand Down Expand Up @@ -144,6 +157,7 @@ export default function usePulseMetadata() {
pulseMetadataList,
pulseMetadataLoading,
loadPulseMetadata,
loadNabatPulseMetadata,
clearPulseMetadata,
viewPulseMetadataLayer,
toggleViewPulseMetadataLayer,
Expand Down
Loading