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