mirror of
https://git.vectorsigma.ru/public/photoview.git
synced 2026-08-03 21:19:18 +00:00
Implement recognizeUnlabeledFaces
This commit is contained in:
@@ -5,6 +5,7 @@ import (
|
||||
|
||||
"github.com/photoview/photoview/api/graphql/auth"
|
||||
"github.com/photoview/photoview/api/graphql/models"
|
||||
"github.com/photoview/photoview/api/scanner/face_detection"
|
||||
"github.com/pkg/errors"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
@@ -119,6 +120,8 @@ func (r *mutationResolver) CombineFaceGroups(ctx context.Context, destinationFac
|
||||
return nil, updateError
|
||||
}
|
||||
|
||||
face_detection.GlobalFaceDetector.MergeCategories(int32(sourceFaceGroupID), int32(destinationFaceGroupID))
|
||||
|
||||
return destinationFaceGroup, nil
|
||||
}
|
||||
|
||||
@@ -127,7 +130,25 @@ func (r *mutationResolver) MoveImageFace(ctx context.Context, imageFaceID int, n
|
||||
}
|
||||
|
||||
func (r *mutationResolver) RecognizeUnlabeledFaces(ctx context.Context) ([]*models.ImageFace, error) {
|
||||
panic("not implemented")
|
||||
user := auth.UserFromContext(ctx)
|
||||
if user == nil {
|
||||
return nil, errors.New("unauthorized")
|
||||
}
|
||||
|
||||
var updatedImageFaces []*models.ImageFace
|
||||
|
||||
transactionError := r.Database.Transaction(func(tx *gorm.DB) error {
|
||||
var err error
|
||||
updatedImageFaces, err = face_detection.GlobalFaceDetector.RecognizeUnlabeledFaces(tx, user)
|
||||
|
||||
return err
|
||||
})
|
||||
|
||||
if transactionError != nil {
|
||||
return nil, transactionError
|
||||
}
|
||||
|
||||
return updatedImageFaces, nil
|
||||
}
|
||||
|
||||
func userOwnedFaceGroup(db *gorm.DB, user *models.User, faceGroupID int) (*models.FaceGroup, error) {
|
||||
|
||||
@@ -103,11 +103,15 @@ func (fd *FaceDetector) DetectFaces(media *models.Media) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (fd *FaceDetector) classifyDescriptor(descriptor face.Descriptor) int32 {
|
||||
return int32(fd.rec.ClassifyThreshold(descriptor, 0.2))
|
||||
}
|
||||
|
||||
func (fd *FaceDetector) classifyFace(face *face.Face, media *models.Media, imagePath string) error {
|
||||
fd.mutex.Lock()
|
||||
defer fd.mutex.Unlock()
|
||||
|
||||
match := fd.rec.ClassifyThreshold(face.Descriptor, 0.2)
|
||||
match := fd.classifyDescriptor(face.Descriptor)
|
||||
|
||||
faceRect, err := models.ToDBFaceRectangle(face.Rectangle, imagePath)
|
||||
if err != nil {
|
||||
@@ -152,3 +156,98 @@ func (fd *FaceDetector) classifyFace(face *face.Face, media *models.Media, image
|
||||
fd.rec.SetSamples(fd.samples, fd.cats)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (fd *FaceDetector) MergeCategories(sourceID int32, destID int32) {
|
||||
fd.mutex.Lock()
|
||||
defer fd.mutex.Unlock()
|
||||
|
||||
for i := range fd.cats {
|
||||
if fd.cats[i] == sourceID {
|
||||
fd.cats[i] = destID
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (fd *FaceDetector) RecognizeUnlabeledFaces(tx *gorm.DB, user *models.User) ([]*models.ImageFace, error) {
|
||||
unrecognizedSamples := make([]face.Descriptor, 0)
|
||||
unrecognizedCats := make([]int32, 0)
|
||||
|
||||
newCats := make([]int32, 0)
|
||||
newSamples := make([]face.Descriptor, 0)
|
||||
|
||||
var unlabeledFaceGroups []*models.FaceGroup
|
||||
|
||||
err := tx.
|
||||
Joins("JOIN image_faces ON image_faces.face_group_id = face_groups.id").
|
||||
Joins("JOIN media ON image_faces.media_id = media.id").
|
||||
Where("face_groups.label IS NULL").
|
||||
Where("media.album_id IN (?)",
|
||||
tx.Select("album_id").Table("user_albums").Where("user_id = ?", user.ID),
|
||||
).
|
||||
Find(&unlabeledFaceGroups).Error
|
||||
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
fd.mutex.Lock()
|
||||
defer fd.mutex.Unlock()
|
||||
|
||||
for i := range fd.samples {
|
||||
cat := fd.cats[i]
|
||||
sample := fd.samples[i]
|
||||
|
||||
catIsUnlabeled := false
|
||||
for _, unlabeledFaceGroup := range unlabeledFaceGroups {
|
||||
if cat == int32(unlabeledFaceGroup.ID) {
|
||||
catIsUnlabeled = true
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
if catIsUnlabeled {
|
||||
unrecognizedCats = append(unrecognizedCats, cat)
|
||||
unrecognizedSamples = append(unrecognizedSamples, sample)
|
||||
} else {
|
||||
newCats = append(newCats, cat)
|
||||
newSamples = append(newSamples, sample)
|
||||
}
|
||||
}
|
||||
|
||||
fd.cats = newCats
|
||||
fd.samples = newSamples
|
||||
|
||||
updatedImageFaces := make([]*models.ImageFace, 0)
|
||||
|
||||
for i := range unrecognizedSamples {
|
||||
cat := unrecognizedCats[i]
|
||||
sample := unrecognizedSamples[i]
|
||||
|
||||
match := fd.classifyDescriptor(sample)
|
||||
|
||||
if match < 0 {
|
||||
// still no match, we can readd it to the list
|
||||
fd.cats = append(fd.cats, cat)
|
||||
fd.samples = append(fd.samples, sample)
|
||||
} else {
|
||||
// found new match, update the database
|
||||
var imageFace models.ImageFace
|
||||
if err := tx.Model(&models.ImageFace{
|
||||
Descriptor: models.FaceDescriptor(sample),
|
||||
}).First(imageFace).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := tx.Model(&imageFace).Update("face_group_id", int(cat)).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
updatedImageFaces = append(updatedImageFaces, &imageFace)
|
||||
|
||||
fd.cats = append(fd.cats, match)
|
||||
fd.samples = append(fd.samples, sample)
|
||||
}
|
||||
}
|
||||
|
||||
return updatedImageFaces, nil
|
||||
}
|
||||
|
||||
@@ -5,7 +5,7 @@ import Layout from '../../Layout'
|
||||
import styled from 'styled-components'
|
||||
import { Link } from 'react-router-dom'
|
||||
import SingleFaceGroup from './SingleFaceGroup/SingleFaceGroup'
|
||||
import { Icon, Input } from 'semantic-ui-react'
|
||||
import { Button, Icon, Input } from 'semantic-ui-react'
|
||||
import FaceCircleImage from './FaceCircleImage'
|
||||
|
||||
export const MY_FACES_QUERY = gql`
|
||||
@@ -43,6 +43,14 @@ export const SET_GROUP_LABEL_MUTATION = gql`
|
||||
}
|
||||
`
|
||||
|
||||
const RECOGNIZE_UNLABELED_FACES_MUTATION = gql`
|
||||
mutation recognizeUnlabeledFaces {
|
||||
recognizeUnlabeledFaces {
|
||||
id
|
||||
}
|
||||
}
|
||||
`
|
||||
|
||||
const FaceDetailsButton = styled.button`
|
||||
color: ${({ labeled }) => (labeled ? 'black' : '#aaa')};
|
||||
margin: 12px auto 24px;
|
||||
@@ -188,6 +196,11 @@ const FaceGroupsWrapper = styled.div`
|
||||
const PeoplePage = ({ match }) => {
|
||||
const { data, error } = useQuery(MY_FACES_QUERY)
|
||||
|
||||
const [
|
||||
recognizeUnlabeled,
|
||||
{ loading: recognizeUnlabeledLoading },
|
||||
] = useMutation(RECOGNIZE_UNLABELED_FACES_MUTATION)
|
||||
|
||||
if (error) {
|
||||
return error.message
|
||||
}
|
||||
@@ -213,6 +226,16 @@ const PeoplePage = ({ match }) => {
|
||||
return (
|
||||
<Layout title={'People'}>
|
||||
<FaceGroupsWrapper>{faces}</FaceGroupsWrapper>
|
||||
<Button
|
||||
loading={recognizeUnlabeledLoading}
|
||||
disabled={recognizeUnlabeledLoading}
|
||||
onClick={() => {
|
||||
recognizeUnlabeled()
|
||||
}}
|
||||
>
|
||||
<Icon name="sync" />
|
||||
Recognize unlabeled faces
|
||||
</Button>
|
||||
</Layout>
|
||||
)
|
||||
}
|
||||
|
||||
@@ -42,7 +42,7 @@ const FaceGroupTitle = ({ faceGroup }) => {
|
||||
)
|
||||
|
||||
const resetLabel = () => {
|
||||
setInputValue(faceGroup.label ?? '')
|
||||
setInputValue(faceGroup?.label ?? '')
|
||||
setEditLabel(false)
|
||||
}
|
||||
|
||||
|
||||
@@ -12,7 +12,7 @@ import FaceCircleImage from '../FaceCircleImage'
|
||||
import { gql, useMutation, useQuery } from '@apollo/client'
|
||||
import { MY_FACES_QUERY } from '../PeoplePage'
|
||||
import styled from 'styled-components'
|
||||
import { Redirect } from 'react-router-dom'
|
||||
import { useHistory } from 'react-router-dom'
|
||||
|
||||
const COMBINE_FACES_MUTATION = gql`
|
||||
mutation($destID: ID!, $srcID: ID!) {
|
||||
@@ -69,13 +69,14 @@ FaceGroupRow.propTypes = {
|
||||
|
||||
const MergeFaceGroupsModal = ({ open, setOpen, sourceFaceGroup }) => {
|
||||
const [page, setPage] = useState(0)
|
||||
const [searchValue, setSearchValue] = useState('')
|
||||
const [selectedRow, setSelectedRow] = useState(null)
|
||||
const [mergedFaceGroup, setMergedFaceGroup] = useState(false)
|
||||
const PAGE_SIZE = 8
|
||||
const PAGE_SIZE = 6
|
||||
let history = useHistory()
|
||||
const { data } = useQuery(MY_FACES_QUERY)
|
||||
const [combineFacesMutation] = useMutation(COMBINE_FACES_MUTATION, {
|
||||
variables: {
|
||||
srcID: sourceFaceGroup.id,
|
||||
srcID: sourceFaceGroup?.id,
|
||||
},
|
||||
refetchQueries: [
|
||||
{
|
||||
@@ -84,27 +85,33 @@ const MergeFaceGroupsModal = ({ open, setOpen, sourceFaceGroup }) => {
|
||||
],
|
||||
})
|
||||
|
||||
if (open == false) return null
|
||||
|
||||
const filteredFaceGroups =
|
||||
data?.myFaceGroups
|
||||
.filter(
|
||||
x =>
|
||||
searchValue == '' ||
|
||||
(x.label && x.label.toLowerCase().includes(searchValue.toLowerCase()))
|
||||
)
|
||||
.filter(x => x.id != sourceFaceGroup?.id) ?? []
|
||||
|
||||
console.log(filteredFaceGroups)
|
||||
|
||||
const mergeFaceGroups = () => {
|
||||
const destFaceGroup = data.myFaceGroups.filter(
|
||||
x => x.id != sourceFaceGroup.id
|
||||
)[selectedRow]
|
||||
const destFaceGroup = filteredFaceGroups[selectedRow]
|
||||
|
||||
combineFacesMutation({
|
||||
variables: {
|
||||
destID: destFaceGroup.id,
|
||||
},
|
||||
onCompleted() {
|
||||
setMergedFaceGroup(destFaceGroup.id)
|
||||
},
|
||||
}).then(() => {
|
||||
setOpen(false)
|
||||
history.push(`/people/${destFaceGroup.id}`)
|
||||
})
|
||||
}
|
||||
|
||||
if (mergedFaceGroup) {
|
||||
return <Redirect to={`/people/${mergedFaceGroup}`} />
|
||||
}
|
||||
|
||||
const rows = data?.myFaceGroups
|
||||
.filter(x => x.id != sourceFaceGroup.id)
|
||||
const rows = filteredFaceGroups
|
||||
.filter((_, i) => i >= page * PAGE_SIZE && i < (page + 1) * PAGE_SIZE)
|
||||
.map((face, i) => (
|
||||
<FaceGroupRow
|
||||
@@ -132,7 +139,13 @@ const MergeFaceGroupsModal = ({ open, setOpen, sourceFaceGroup }) => {
|
||||
</Table.Row>
|
||||
<Table.Row>
|
||||
<Table.HeaderCell>
|
||||
<Input icon="search" placeholder="Search faces..." fluid />
|
||||
<Input
|
||||
value={searchValue}
|
||||
onChange={e => setSearchValue(e.target.value)}
|
||||
icon="search"
|
||||
placeholder="Search faces..."
|
||||
fluid
|
||||
/>
|
||||
</Table.HeaderCell>
|
||||
</Table.Row>
|
||||
</Table.Header>
|
||||
@@ -176,7 +189,7 @@ const MergeFaceGroupsModal = ({ open, setOpen, sourceFaceGroup }) => {
|
||||
MergeFaceGroupsModal.propTypes = {
|
||||
open: PropTypes.bool.isRequired,
|
||||
setOpen: PropTypes.func.isRequired,
|
||||
sourceFaceGroup: PropTypes.object.isRequired,
|
||||
sourceFaceGroup: PropTypes.object,
|
||||
}
|
||||
|
||||
export default MergeFaceGroupsModal
|
||||
|
||||
Reference in New Issue
Block a user