diff --git a/api/graphql/models/face_detection.go b/api/graphql/models/face_detection.go index 8e09d037..32e98db3 100644 --- a/api/graphql/models/face_detection.go +++ b/api/graphql/models/face_detection.go @@ -9,7 +9,6 @@ import ( "strconv" "strings" - "github.com/Kagami/go-face" "github.com/photoview/photoview/api/database/drivers" "github.com/photoview/photoview/api/scanner/media_encoding/media_utils" "gorm.io/gorm" @@ -45,7 +44,7 @@ func (f *ImageFace) FillMedia(db *gorm.DB) error { return nil } -type FaceDescriptor face.Descriptor +type FaceDescriptor [128]float32 // same as go-face's Descriptor // GormDataType datatype used in database func (FaceDescriptor) GormDBDataType(db *gorm.DB, field *schema.Field) string { diff --git a/api/scanner/face_detection/face_detector.go b/api/scanner/face_detection/face_detector.go index 506cb887..15a69eca 100644 --- a/api/scanner/face_detection/face_detector.go +++ b/api/scanner/face_detection/face_detector.go @@ -1,300 +1,17 @@ package face_detection import ( - "log" - "sync" - - "github.com/Kagami/go-face" - "github.com/photoview/photoview/api/graphql/models" - "github.com/photoview/photoview/api/utils" - "github.com/pkg/errors" "gorm.io/gorm" + + "github.com/photoview/photoview/api/graphql/models" ) -type FaceDetector struct { - mutex sync.Mutex - rec *face.Recognizer - faceDescriptors []face.Descriptor - faceGroupIDs []int32 - imageFaceIDs []int +type FaceDetector interface { + ReloadFacesFromDatabase(db *gorm.DB) error + DetectFaces(db *gorm.DB, media *models.Media) error + MergeCategories(sourceID int32, destID int32) + MergeImageFaces(imageFaceIDs []int, destFaceGroupID int32) + RecognizeUnlabeledFaces(tx *gorm.DB, user *models.User) ([]*models.ImageFace, error) } -var GlobalFaceDetector *FaceDetector = nil - -func InitializeFaceDetector(db *gorm.DB) error { - if utils.EnvDisableFaceRecognition.GetBool() { - log.Printf("Face detection disabled (%s=1)\n", utils.EnvDisableFaceRecognition.GetName()) - return nil - } - - log.Println("Initializing face detector") - - rec, err := face.NewRecognizer(utils.FaceRecognitionModelsPath()) - if err != nil { - return errors.Wrap(err, "initialize facedetect recognizer") - } - - faceDescriptors, faceGroupIDs, imageFaceIDs, err := getSamplesFromDatabase(db) - if err != nil { - return errors.Wrap(err, "get face detection samples from database") - } - - GlobalFaceDetector = &FaceDetector{ - rec: rec, - faceDescriptors: faceDescriptors, - faceGroupIDs: faceGroupIDs, - imageFaceIDs: imageFaceIDs, - } - - return nil -} - -func getSamplesFromDatabase(db *gorm.DB) (samples []face.Descriptor, faceGroupIDs []int32, imageFaceIDs []int, err error) { - - var imageFaces []*models.ImageFace - - if err = db.Find(&imageFaces).Error; err != nil { - return - } - - samples = make([]face.Descriptor, len(imageFaces)) - faceGroupIDs = make([]int32, len(imageFaces)) - imageFaceIDs = make([]int, len(imageFaces)) - - for i, imgFace := range imageFaces { - samples[i] = face.Descriptor(imgFace.Descriptor) - faceGroupIDs[i] = int32(imgFace.FaceGroupID) - imageFaceIDs[i] = imgFace.ID - } - - return -} - -// ReloadFacesFromDatabase replaces the in-memory face descriptors with the ones in the database -func (fd *FaceDetector) ReloadFacesFromDatabase(db *gorm.DB) error { - faceDescriptors, faceGroupIDs, imageFaceIDs, err := getSamplesFromDatabase(db) - if err != nil { - return err - } - - fd.mutex.Lock() - defer fd.mutex.Unlock() - - fd.faceDescriptors = faceDescriptors - fd.faceGroupIDs = faceGroupIDs - fd.imageFaceIDs = imageFaceIDs - - return nil -} - -// DetectFaces finds the faces in the given image and saves them to the database -func (fd *FaceDetector) DetectFaces(db *gorm.DB, media *models.Media) error { - if err := db.Model(media).Preload("MediaURL").First(&media).Error; err != nil { - return err - } - - var thumbnailURL *models.MediaURL - for _, url := range media.MediaURL { - if url.Purpose == models.PhotoThumbnail { - thumbnailURL = &url - thumbnailURL.Media = media - break - } - } - - if thumbnailURL == nil { - return errors.New("thumbnail url is missing") - } - - thumbnailPath, err := thumbnailURL.CachedPath() - if err != nil { - return err - } - - fd.mutex.Lock() - faces, err := fd.rec.RecognizeFile(thumbnailPath) - fd.mutex.Unlock() - - if err != nil { - return errors.Wrap(err, "error read faces") - } - - for _, face := range faces { - fd.classifyFace(db, &face, media, thumbnailPath) - } - - return nil -} - -func (fd *FaceDetector) classifyDescriptor(descriptor face.Descriptor) int32 { - return int32(fd.rec.ClassifyThreshold(descriptor, 0.2)) -} - -func (fd *FaceDetector) classifyFace(db *gorm.DB, face *face.Face, media *models.Media, imagePath string) error { - fd.mutex.Lock() - defer fd.mutex.Unlock() - - match := fd.classifyDescriptor(face.Descriptor) - - faceRect, err := models.ToDBFaceRectangle(face.Rectangle, imagePath) - if err != nil { - return err - } - - imageFace := models.ImageFace{ - MediaID: media.ID, - Descriptor: models.FaceDescriptor(face.Descriptor), - Rectangle: *faceRect, - } - - var faceGroup models.FaceGroup - - // If no match add it new to samples - if match < 0 { - log.Println("No match, assigning new face") - - faceGroup = models.FaceGroup{ - ImageFaces: []models.ImageFace{imageFace}, - } - - if err := db.Create(&faceGroup).Error; err != nil { - return err - } - - } else { - log.Println("Found match") - - if err := db.First(&faceGroup, int(match)).Error; err != nil { - return err - } - - if err := db.Model(&faceGroup).Association("ImageFaces").Append(&imageFace); err != nil { - return err - } - } - - fd.faceDescriptors = append(fd.faceDescriptors, face.Descriptor) - fd.faceGroupIDs = append(fd.faceGroupIDs, int32(faceGroup.ID)) - fd.imageFaceIDs = append(fd.imageFaceIDs, imageFace.ID) - - fd.rec.SetSamples(fd.faceDescriptors, fd.faceGroupIDs) - return nil -} - -func (fd *FaceDetector) MergeCategories(sourceID int32, destID int32) { - fd.mutex.Lock() - defer fd.mutex.Unlock() - - for i := range fd.faceGroupIDs { - if fd.faceGroupIDs[i] == sourceID { - fd.faceGroupIDs[i] = destID - } - } -} - -func (fd *FaceDetector) MergeImageFaces(imageFaceIDs []int, destFaceGroupID int32) { - fd.mutex.Lock() - defer fd.mutex.Unlock() - - for i := range fd.faceGroupIDs { - imageFaceID := fd.imageFaceIDs[i] - - for _, id := range imageFaceIDs { - if imageFaceID == id { - fd.faceGroupIDs[i] = destFaceGroupID - break - } - } - } -} - -func (fd *FaceDetector) RecognizeUnlabeledFaces(tx *gorm.DB, user *models.User) ([]*models.ImageFace, error) { - unrecognizedDescriptors := make([]face.Descriptor, 0) - unrecognizedFaceGroupIDs := make([]int32, 0) - unrecognizedImageFaceIDs := make([]int, 0) - - newFaceGroupIDs := make([]int32, 0) - newDescriptors := make([]face.Descriptor, 0) - newImageFaceIDs := make([]int, 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.faceDescriptors { - descriptor := fd.faceDescriptors[i] - faceGroupID := fd.faceGroupIDs[i] - imageFaceID := fd.imageFaceIDs[i] - - isUnlabeled := false - for _, unlabeledFaceGroup := range unlabeledFaceGroups { - if faceGroupID == int32(unlabeledFaceGroup.ID) { - isUnlabeled = true - continue - } - } - - if isUnlabeled { - unrecognizedFaceGroupIDs = append(unrecognizedFaceGroupIDs, faceGroupID) - unrecognizedDescriptors = append(unrecognizedDescriptors, descriptor) - unrecognizedImageFaceIDs = append(unrecognizedImageFaceIDs, imageFaceID) - } else { - newFaceGroupIDs = append(newFaceGroupIDs, faceGroupID) - newDescriptors = append(newDescriptors, descriptor) - newImageFaceIDs = append(newImageFaceIDs, imageFaceID) - } - } - - fd.faceGroupIDs = newFaceGroupIDs - fd.faceDescriptors = newDescriptors - fd.imageFaceIDs = newImageFaceIDs - - updatedImageFaces := make([]*models.ImageFace, 0) - - for i := range unrecognizedDescriptors { - descriptor := unrecognizedDescriptors[i] - faceGroupID := unrecognizedFaceGroupIDs[i] - imageFaceID := unrecognizedImageFaceIDs[i] - - match := fd.classifyDescriptor(descriptor) - - if match < 0 { - // still no match, we can readd it to the list - fd.faceGroupIDs = append(fd.faceGroupIDs, faceGroupID) - fd.faceDescriptors = append(fd.faceDescriptors, descriptor) - fd.imageFaceIDs = append(fd.imageFaceIDs, imageFaceID) - } else { - // found new match, update the database - var imageFace models.ImageFace - if err := tx.Model(&models.ImageFace{}).First(imageFace, imageFaceID).Error; err != nil { - return nil, err - } - - if err := tx.Model(&imageFace).Update("face_group_id", int(faceGroupID)).Error; err != nil { - return nil, err - } - - updatedImageFaces = append(updatedImageFaces, &imageFace) - - fd.faceGroupIDs = append(fd.faceGroupIDs, match) - fd.faceDescriptors = append(fd.faceDescriptors, descriptor) - fd.imageFaceIDs = append(fd.imageFaceIDs, imageFaceID) - } - } - - return updatedImageFaces, nil -} +var GlobalFaceDetector FaceDetector = nil diff --git a/api/scanner/face_detection/face_detector_impl.go b/api/scanner/face_detection/face_detector_impl.go new file mode 100644 index 00000000..e28aa0c3 --- /dev/null +++ b/api/scanner/face_detection/face_detector_impl.go @@ -0,0 +1,300 @@ +//go:build !no_face_detection + +package face_detection + +import ( + "log" + "sync" + + "github.com/Kagami/go-face" + "github.com/photoview/photoview/api/graphql/models" + "github.com/photoview/photoview/api/utils" + "github.com/pkg/errors" + "gorm.io/gorm" +) + +type faceDetector struct { + mutex sync.Mutex + rec *face.Recognizer + faceDescriptors []face.Descriptor + faceGroupIDs []int32 + imageFaceIDs []int +} + +func InitializeFaceDetector(db *gorm.DB) error { + if utils.EnvDisableFaceRecognition.GetBool() { + log.Printf("Face detection disabled (%s=1)\n", utils.EnvDisableFaceRecognition.GetName()) + return nil + } + + log.Println("Initializing face detector") + + rec, err := face.NewRecognizer(utils.FaceRecognitionModelsPath()) + if err != nil { + return errors.Wrap(err, "initialize facedetect recognizer") + } + + faceDescriptors, faceGroupIDs, imageFaceIDs, err := getSamplesFromDatabase(db) + if err != nil { + return errors.Wrap(err, "get face detection samples from database") + } + + GlobalFaceDetector = &faceDetector{ + rec: rec, + faceDescriptors: faceDescriptors, + faceGroupIDs: faceGroupIDs, + imageFaceIDs: imageFaceIDs, + } + + return nil +} + +func getSamplesFromDatabase(db *gorm.DB) (samples []face.Descriptor, faceGroupIDs []int32, imageFaceIDs []int, err error) { + + var imageFaces []*models.ImageFace + + if err = db.Find(&imageFaces).Error; err != nil { + return + } + + samples = make([]face.Descriptor, len(imageFaces)) + faceGroupIDs = make([]int32, len(imageFaces)) + imageFaceIDs = make([]int, len(imageFaces)) + + for i, imgFace := range imageFaces { + samples[i] = face.Descriptor(imgFace.Descriptor) + faceGroupIDs[i] = int32(imgFace.FaceGroupID) + imageFaceIDs[i] = imgFace.ID + } + + return +} + +// ReloadFacesFromDatabase replaces the in-memory face descriptors with the ones in the database +func (fd *faceDetector) ReloadFacesFromDatabase(db *gorm.DB) error { + faceDescriptors, faceGroupIDs, imageFaceIDs, err := getSamplesFromDatabase(db) + if err != nil { + return err + } + + fd.mutex.Lock() + defer fd.mutex.Unlock() + + fd.faceDescriptors = faceDescriptors + fd.faceGroupIDs = faceGroupIDs + fd.imageFaceIDs = imageFaceIDs + + return nil +} + +// DetectFaces finds the faces in the given image and saves them to the database +func (fd *faceDetector) DetectFaces(db *gorm.DB, media *models.Media) error { + if err := db.Model(media).Preload("MediaURL").First(&media).Error; err != nil { + return err + } + + var thumbnailURL *models.MediaURL + for _, url := range media.MediaURL { + if url.Purpose == models.PhotoThumbnail { + thumbnailURL = &url + thumbnailURL.Media = media + break + } + } + + if thumbnailURL == nil { + return errors.New("thumbnail url is missing") + } + + thumbnailPath, err := thumbnailURL.CachedPath() + if err != nil { + return err + } + + fd.mutex.Lock() + faces, err := fd.rec.RecognizeFile(thumbnailPath) + fd.mutex.Unlock() + + if err != nil { + return errors.Wrap(err, "error read faces") + } + + for _, face := range faces { + fd.classifyFace(db, &face, media, thumbnailPath) + } + + return nil +} + +func (fd *faceDetector) classifyDescriptor(descriptor face.Descriptor) int32 { + return int32(fd.rec.ClassifyThreshold(descriptor, 0.2)) +} + +func (fd *faceDetector) classifyFace(db *gorm.DB, face *face.Face, media *models.Media, imagePath string) error { + fd.mutex.Lock() + defer fd.mutex.Unlock() + + match := fd.classifyDescriptor(face.Descriptor) + + faceRect, err := models.ToDBFaceRectangle(face.Rectangle, imagePath) + if err != nil { + return err + } + + imageFace := models.ImageFace{ + MediaID: media.ID, + Descriptor: models.FaceDescriptor(face.Descriptor), + Rectangle: *faceRect, + } + + var faceGroup models.FaceGroup + + // If no match add it new to samples + if match < 0 { + log.Println("No match, assigning new face") + + faceGroup = models.FaceGroup{ + ImageFaces: []models.ImageFace{imageFace}, + } + + if err := db.Create(&faceGroup).Error; err != nil { + return err + } + + } else { + log.Println("Found match") + + if err := db.First(&faceGroup, int(match)).Error; err != nil { + return err + } + + if err := db.Model(&faceGroup).Association("ImageFaces").Append(&imageFace); err != nil { + return err + } + } + + fd.faceDescriptors = append(fd.faceDescriptors, face.Descriptor) + fd.faceGroupIDs = append(fd.faceGroupIDs, int32(faceGroup.ID)) + fd.imageFaceIDs = append(fd.imageFaceIDs, imageFace.ID) + + fd.rec.SetSamples(fd.faceDescriptors, fd.faceGroupIDs) + return nil +} + +func (fd *faceDetector) MergeCategories(sourceID int32, destID int32) { + fd.mutex.Lock() + defer fd.mutex.Unlock() + + for i := range fd.faceGroupIDs { + if fd.faceGroupIDs[i] == sourceID { + fd.faceGroupIDs[i] = destID + } + } +} + +func (fd *faceDetector) MergeImageFaces(imageFaceIDs []int, destFaceGroupID int32) { + fd.mutex.Lock() + defer fd.mutex.Unlock() + + for i := range fd.faceGroupIDs { + imageFaceID := fd.imageFaceIDs[i] + + for _, id := range imageFaceIDs { + if imageFaceID == id { + fd.faceGroupIDs[i] = destFaceGroupID + break + } + } + } +} + +func (fd *faceDetector) RecognizeUnlabeledFaces(tx *gorm.DB, user *models.User) ([]*models.ImageFace, error) { + unrecognizedDescriptors := make([]face.Descriptor, 0) + unrecognizedFaceGroupIDs := make([]int32, 0) + unrecognizedImageFaceIDs := make([]int, 0) + + newFaceGroupIDs := make([]int32, 0) + newDescriptors := make([]face.Descriptor, 0) + newImageFaceIDs := make([]int, 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.faceDescriptors { + descriptor := fd.faceDescriptors[i] + faceGroupID := fd.faceGroupIDs[i] + imageFaceID := fd.imageFaceIDs[i] + + isUnlabeled := false + for _, unlabeledFaceGroup := range unlabeledFaceGroups { + if faceGroupID == int32(unlabeledFaceGroup.ID) { + isUnlabeled = true + continue + } + } + + if isUnlabeled { + unrecognizedFaceGroupIDs = append(unrecognizedFaceGroupIDs, faceGroupID) + unrecognizedDescriptors = append(unrecognizedDescriptors, descriptor) + unrecognizedImageFaceIDs = append(unrecognizedImageFaceIDs, imageFaceID) + } else { + newFaceGroupIDs = append(newFaceGroupIDs, faceGroupID) + newDescriptors = append(newDescriptors, descriptor) + newImageFaceIDs = append(newImageFaceIDs, imageFaceID) + } + } + + fd.faceGroupIDs = newFaceGroupIDs + fd.faceDescriptors = newDescriptors + fd.imageFaceIDs = newImageFaceIDs + + updatedImageFaces := make([]*models.ImageFace, 0) + + for i := range unrecognizedDescriptors { + descriptor := unrecognizedDescriptors[i] + faceGroupID := unrecognizedFaceGroupIDs[i] + imageFaceID := unrecognizedImageFaceIDs[i] + + match := fd.classifyDescriptor(descriptor) + + if match < 0 { + // still no match, we can readd it to the list + fd.faceGroupIDs = append(fd.faceGroupIDs, faceGroupID) + fd.faceDescriptors = append(fd.faceDescriptors, descriptor) + fd.imageFaceIDs = append(fd.imageFaceIDs, imageFaceID) + } else { + // found new match, update the database + var imageFace models.ImageFace + if err := tx.Model(&models.ImageFace{}).First(imageFace, imageFaceID).Error; err != nil { + return nil, err + } + + if err := tx.Model(&imageFace).Update("face_group_id", int(faceGroupID)).Error; err != nil { + return nil, err + } + + updatedImageFaces = append(updatedImageFaces, &imageFace) + + fd.faceGroupIDs = append(fd.faceGroupIDs, match) + fd.faceDescriptors = append(fd.faceDescriptors, descriptor) + fd.imageFaceIDs = append(fd.imageFaceIDs, imageFaceID) + } + } + + return updatedImageFaces, nil +} diff --git a/api/scanner/face_detection/face_detector_shim.go b/api/scanner/face_detection/face_detector_shim.go new file mode 100644 index 00000000..54906952 --- /dev/null +++ b/api/scanner/face_detection/face_detector_shim.go @@ -0,0 +1,14 @@ +//go:build no_face_detection + +package face_detection + +import ( + "log" + + "gorm.io/gorm" +) + +func InitializeFaceDetector(db *gorm.DB) error { + log.Printf("Face detection disabled (at build-time)") + return nil +}