Implement recognizeUnlabeledFaces

This commit is contained in:
viktorstrate
2021-02-19 19:24:31 +01:00
parent 00fceea4db
commit bdd2318afc
5 changed files with 178 additions and 22 deletions

View File

@@ -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) {

View File

@@ -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
}

View File

@@ -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>
)
}

View File

@@ -42,7 +42,7 @@ const FaceGroupTitle = ({ faceGroup }) => {
)
const resetLabel = () => {
setInputValue(faceGroup.label ?? '')
setInputValue(faceGroup?.label ?? '')
setEditLabel(false)
}

View File

@@ -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