From bdd2318afcaa07c9b7f674369987ce4e7367fefd Mon Sep 17 00:00:00 2001 From: viktorstrate Date: Fri, 19 Feb 2021 19:24:31 +0100 Subject: [PATCH] Implement recognizeUnlabeledFaces --- api/graphql/resolvers/faces.go | 23 +++- api/scanner/face_detection/face_detector.go | 101 +++++++++++++++++- ui/src/Pages/PeoplePage/PeoplePage.js | 25 ++++- .../SingleFaceGroup/FaceGroupTitle.js | 2 +- .../SingleFaceGroup/MergeFaceGroupsModal.js | 49 +++++---- 5 files changed, 178 insertions(+), 22 deletions(-) diff --git a/api/graphql/resolvers/faces.go b/api/graphql/resolvers/faces.go index 885de7b3..fa4ff040 100644 --- a/api/graphql/resolvers/faces.go +++ b/api/graphql/resolvers/faces.go @@ -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) { diff --git a/api/scanner/face_detection/face_detector.go b/api/scanner/face_detection/face_detector.go index 5bbca944..9cab8885 100644 --- a/api/scanner/face_detection/face_detector.go +++ b/api/scanner/face_detection/face_detector.go @@ -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 +} diff --git a/ui/src/Pages/PeoplePage/PeoplePage.js b/ui/src/Pages/PeoplePage/PeoplePage.js index d64e5455..8bbfc391 100644 --- a/ui/src/Pages/PeoplePage/PeoplePage.js +++ b/ui/src/Pages/PeoplePage/PeoplePage.js @@ -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 ( {faces} + ) } diff --git a/ui/src/Pages/PeoplePage/SingleFaceGroup/FaceGroupTitle.js b/ui/src/Pages/PeoplePage/SingleFaceGroup/FaceGroupTitle.js index edf71789..6836b2b8 100644 --- a/ui/src/Pages/PeoplePage/SingleFaceGroup/FaceGroupTitle.js +++ b/ui/src/Pages/PeoplePage/SingleFaceGroup/FaceGroupTitle.js @@ -42,7 +42,7 @@ const FaceGroupTitle = ({ faceGroup }) => { ) const resetLabel = () => { - setInputValue(faceGroup.label ?? '') + setInputValue(faceGroup?.label ?? '') setEditLabel(false) } diff --git a/ui/src/Pages/PeoplePage/SingleFaceGroup/MergeFaceGroupsModal.js b/ui/src/Pages/PeoplePage/SingleFaceGroup/MergeFaceGroupsModal.js index 5742f98a..c38daaf6 100644 --- a/ui/src/Pages/PeoplePage/SingleFaceGroup/MergeFaceGroupsModal.js +++ b/ui/src/Pages/PeoplePage/SingleFaceGroup/MergeFaceGroupsModal.js @@ -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 - } - - 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) => ( { - + setSearchValue(e.target.value)} + icon="search" + placeholder="Search faces..." + fluid + /> @@ -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