Extract detection and segmentation workflow hooks
This commit is contained in:
@@ -0,0 +1,227 @@
|
||||
import { useMemo, useState } from 'react'
|
||||
import { segmentationApi } from '../services/api'
|
||||
import type {
|
||||
DatasetCreateResponse,
|
||||
QualityCheckRead,
|
||||
SegmentationModelCapability,
|
||||
SegmentationQaResult,
|
||||
SegmentationRead,
|
||||
SegmentationRunRead,
|
||||
SegmentationRunResponse,
|
||||
} from '../types'
|
||||
import { formatError } from '../lib/formatError'
|
||||
|
||||
interface SegmentationWorkflowOptions {
|
||||
selectedProjectId: string | null
|
||||
rasterDatasets: DatasetCreateResponse[]
|
||||
qaIouThreshold: number
|
||||
loadProjectData: (projectId: string) => Promise<unknown>
|
||||
loadQualityChecks: (projectId?: string | null) => Promise<QualityCheckRead[] | void>
|
||||
}
|
||||
|
||||
export function useSegmentationWorkflow({
|
||||
selectedProjectId,
|
||||
rasterDatasets,
|
||||
qaIouThreshold,
|
||||
loadProjectData,
|
||||
loadQualityChecks,
|
||||
}: SegmentationWorkflowOptions) {
|
||||
const [segmentationModels, setSegmentationModels] = useState<SegmentationModelCapability[]>([])
|
||||
const [loadingSegmentationModels, setLoadingSegmentationModels] = useState(false)
|
||||
const [segmentationModelError, setSegmentationModelError] = useState<string | null>(null)
|
||||
const [selectedSegmentationDatasetId, setSelectedSegmentationDatasetId] = useState('')
|
||||
const [selectedSegmentationModelId, setSelectedSegmentationModelId] = useState('segmentation-placeholder')
|
||||
const [segmentationConfidenceThreshold, setSegmentationConfidenceThreshold] = useState(0.5)
|
||||
const [runningSegmentation, setRunningSegmentation] = useState(false)
|
||||
const [segmentationRunResult, setSegmentationRunResult] = useState<SegmentationRunResponse | null>(null)
|
||||
const [segmentationRunError, setSegmentationRunError] = useState<string | null>(null)
|
||||
const [segmentationRuns, setSegmentationRuns] = useState<SegmentationRunRead[]>([])
|
||||
const [selectedSegmentationRunId, setSelectedSegmentationRunId] = useState('')
|
||||
const [segmentationItems, setSegmentationItems] = useState<SegmentationRead[]>([])
|
||||
const [segmentationGeoJson, setSegmentationGeoJson] = useState<GeoJSON.FeatureCollection | null>(null)
|
||||
const [segmentationClassFilter, setSegmentationClassFilter] = useState('')
|
||||
const [segmentationMinConfidenceFilter, setSegmentationMinConfidenceFilter] = useState(0)
|
||||
const [loadingSegmentationResults, setLoadingSegmentationResults] = useState(false)
|
||||
const [segmentationReferenceDatasetId, setSegmentationReferenceDatasetId] = useState('')
|
||||
const [segmentationQaResult, setSegmentationQaResult] = useState<SegmentationQaResult | null>(null)
|
||||
const [segmentationQaError, setSegmentationQaError] = useState<string | null>(null)
|
||||
const [runningSegmentationQa, setRunningSegmentationQa] = useState(false)
|
||||
|
||||
const selectedSegmentationModel = useMemo(
|
||||
() => segmentationModels.find((model) => model.model_id === selectedSegmentationModelId) ?? null,
|
||||
[segmentationModels, selectedSegmentationModelId],
|
||||
)
|
||||
|
||||
const loadSegmentationModels = async () => {
|
||||
setLoadingSegmentationModels(true)
|
||||
setSegmentationModelError(null)
|
||||
try {
|
||||
const response = await segmentationApi.listModels()
|
||||
setSegmentationModels(response.models)
|
||||
if (!response.models.some((model) => model.model_id === selectedSegmentationModelId) && response.models.length > 0) {
|
||||
setSelectedSegmentationModelId(response.models[0].model_id)
|
||||
}
|
||||
} catch (error) {
|
||||
setSegmentationModelError(formatError(error, 'Failed to load segmentation models'))
|
||||
} finally {
|
||||
setLoadingSegmentationModels(false)
|
||||
}
|
||||
}
|
||||
|
||||
const loadSegmentationRuns = async (projectId = selectedProjectId) => {
|
||||
if (!projectId) {
|
||||
setSegmentationRuns([])
|
||||
return
|
||||
}
|
||||
try {
|
||||
const response = await segmentationApi.listRuns({ project_id: projectId })
|
||||
setSegmentationRuns(response.items)
|
||||
if (!selectedSegmentationRunId && response.items.length > 0) {
|
||||
setSelectedSegmentationRunId(response.items[0].id)
|
||||
}
|
||||
} catch (error) {
|
||||
setSegmentationRunError(formatError(error, 'Failed to load segmentation runs'))
|
||||
}
|
||||
}
|
||||
|
||||
const loadSegmentationResults = async (analysisRunId = selectedSegmentationRunId) => {
|
||||
if (!analysisRunId) {
|
||||
setSegmentationItems([])
|
||||
setSegmentationGeoJson(null)
|
||||
return
|
||||
}
|
||||
setLoadingSegmentationResults(true)
|
||||
setSegmentationRunError(null)
|
||||
try {
|
||||
const params = {
|
||||
class_name: segmentationClassFilter || null,
|
||||
min_confidence: segmentationMinConfidenceFilter > 0 ? segmentationMinConfidenceFilter : null,
|
||||
}
|
||||
const [segmentationsResponse, geoJsonResponse] = await Promise.all([
|
||||
segmentationApi.listSegmentations(analysisRunId, params),
|
||||
segmentationApi.getRunGeoJson(analysisRunId, params),
|
||||
])
|
||||
setSegmentationItems(segmentationsResponse.items)
|
||||
setSegmentationGeoJson(geoJsonResponse)
|
||||
} catch (error) {
|
||||
setSegmentationRunError(formatError(error, 'Failed to load segmentation results'))
|
||||
} finally {
|
||||
setLoadingSegmentationResults(false)
|
||||
}
|
||||
}
|
||||
|
||||
const runSegmentation = async () => {
|
||||
if (!selectedProjectId) {
|
||||
setSegmentationRunError('Select a project first')
|
||||
return
|
||||
}
|
||||
const datasetId = selectedSegmentationDatasetId || rasterDatasets[0]?.id
|
||||
if (!datasetId) {
|
||||
setSegmentationRunError('Select a raster dataset')
|
||||
return
|
||||
}
|
||||
if (!selectedSegmentationModel?.configured) {
|
||||
setSegmentationRunError('Selected segmentation model is not configured')
|
||||
return
|
||||
}
|
||||
setSegmentationRunError(null)
|
||||
setSegmentationRunResult(null)
|
||||
setRunningSegmentation(true)
|
||||
try {
|
||||
const parameters =
|
||||
selectedSegmentationModelId === 'fixture-segmenter'
|
||||
? { fixture_mode: true, fixture_segmentations: [] }
|
||||
: {}
|
||||
const result = await segmentationApi.run({
|
||||
project_id: selectedProjectId,
|
||||
dataset_id: datasetId,
|
||||
model_id: selectedSegmentationModelId,
|
||||
confidence_threshold: segmentationConfidenceThreshold,
|
||||
parameters_json: parameters,
|
||||
})
|
||||
setSegmentationRunResult(result)
|
||||
setSelectedSegmentationRunId(result.analysis_run_id)
|
||||
await loadSegmentationRuns(selectedProjectId)
|
||||
await loadSegmentationResults(result.analysis_run_id)
|
||||
await loadProjectData(selectedProjectId)
|
||||
} catch (error) {
|
||||
setSegmentationRunError(formatError(error, 'Segmentation run failed'))
|
||||
} finally {
|
||||
setRunningSegmentation(false)
|
||||
}
|
||||
}
|
||||
|
||||
const runSegmentationQa = async () => {
|
||||
if (!selectedSegmentationRunId) {
|
||||
setSegmentationQaError('Select a segmentation run')
|
||||
return
|
||||
}
|
||||
if (!segmentationReferenceDatasetId) {
|
||||
setSegmentationQaError('Select a reference dataset')
|
||||
return
|
||||
}
|
||||
setSegmentationQaError(null)
|
||||
setSegmentationQaResult(null)
|
||||
setRunningSegmentationQa(true)
|
||||
try {
|
||||
const result = await segmentationApi.compareWithReference(selectedSegmentationRunId, {
|
||||
reference_dataset_id: segmentationReferenceDatasetId,
|
||||
iou_threshold: qaIouThreshold,
|
||||
class_name: segmentationClassFilter || null,
|
||||
min_confidence: segmentationMinConfidenceFilter > 0 ? segmentationMinConfidenceFilter : null,
|
||||
})
|
||||
setSegmentationQaResult(result)
|
||||
await loadQualityChecks(selectedProjectId)
|
||||
} catch (error) {
|
||||
setSegmentationQaError(formatError(error, 'Segmentation QA failed'))
|
||||
} finally {
|
||||
setRunningSegmentationQa(false)
|
||||
}
|
||||
}
|
||||
|
||||
const resetSegmentationForProject = () => {
|
||||
setSelectedSegmentationDatasetId('')
|
||||
setSegmentationRuns([])
|
||||
setSelectedSegmentationRunId('')
|
||||
setSegmentationItems([])
|
||||
setSegmentationGeoJson(null)
|
||||
setSegmentationRunResult(null)
|
||||
}
|
||||
|
||||
return {
|
||||
segmentationModels,
|
||||
loadingSegmentationModels,
|
||||
segmentationModelError,
|
||||
selectedSegmentationDatasetId,
|
||||
selectedSegmentationModelId,
|
||||
selectedSegmentationModel,
|
||||
segmentationConfidenceThreshold,
|
||||
runningSegmentation,
|
||||
segmentationRunResult,
|
||||
segmentationRunError,
|
||||
segmentationRuns,
|
||||
selectedSegmentationRunId,
|
||||
segmentationItems,
|
||||
segmentationGeoJson,
|
||||
segmentationClassFilter,
|
||||
segmentationMinConfidenceFilter,
|
||||
loadingSegmentationResults,
|
||||
segmentationReferenceDatasetId,
|
||||
segmentationQaResult,
|
||||
segmentationQaError,
|
||||
runningSegmentationQa,
|
||||
loadSegmentationModels,
|
||||
loadSegmentationRuns,
|
||||
loadSegmentationResults,
|
||||
runSegmentation,
|
||||
runSegmentationQa,
|
||||
resetSegmentationForProject,
|
||||
setSelectedSegmentationDatasetId,
|
||||
setSelectedSegmentationModelId,
|
||||
setSegmentationConfidenceThreshold,
|
||||
setSelectedSegmentationRunId,
|
||||
setSegmentationClassFilter,
|
||||
setSegmentationMinConfidenceFilter,
|
||||
setSegmentationReferenceDatasetId,
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user