Add UI to move faces between face groups

This commit is contained in:
viktorstrate
2021-02-20 22:43:07 +01:00
parent a3e5346501
commit 20251dedd6
11 changed files with 611 additions and 189 deletions

View File

@@ -150,7 +150,83 @@ func (r *mutationResolver) CombineFaceGroups(ctx context.Context, destinationFac
}
func (r *mutationResolver) MoveImageFaces(ctx context.Context, imageFaceIDs []int, destinationFaceGroupID int) (*models.FaceGroup, error) {
panic("not implemented")
user := auth.UserFromContext(ctx)
if user == nil {
return nil, errors.New("unauthorized")
}
if err := user.FillAlbums(r.Database); err != nil {
return nil, err
}
userAlbumIDs := make([]int, len(user.Albums))
for i, album := range user.Albums {
userAlbumIDs[i] = album.ID
}
userOwnedImageFaceIDs := make([]int, 0)
var destFaceGroup *models.FaceGroup
transErr := r.Database.Transaction(func(tx *gorm.DB) error {
var err error
destFaceGroup, err = userOwnedFaceGroup(tx, user, destinationFaceGroupID)
if err != nil {
return err
}
var userOwnedImageFaces []*models.ImageFace
if err := tx.
Joins("JOIN media ON media.id = image_faces.media_id").
Where("media.album_id IN (?)", userAlbumIDs).
Where("image_faces.id IN (?)", imageFaceIDs).
Find(&userOwnedImageFaces).Error; err != nil {
return err
}
for _, imageFace := range userOwnedImageFaces {
userOwnedImageFaceIDs = append(userOwnedImageFaceIDs, imageFace.ID)
}
var sourceFaceGroups []*models.FaceGroup
if err := tx.
Joins("LEFT JOIN image_faces ON image_faces.face_group_id = face_groups.id").
Where("image_faces.id IN (?)", userOwnedImageFaceIDs).
Find(&sourceFaceGroups).Error; err != nil {
return err
}
if err := tx.
Model(&models.ImageFace{}).
Where("id IN (?)", userOwnedImageFaceIDs).
Update("face_group_id", destFaceGroup.ID).Error; err != nil {
return err
}
// delete face groups if they have become empty
for _, faceGroup := range sourceFaceGroups {
var count int64
if err := tx.Model(&models.ImageFace{}).Where("face_group_id = ?", faceGroup.ID).Count(&count).Error; err != nil {
return err
}
if count == 0 {
if err := tx.Delete(&faceGroup).Error; err != nil {
return err
}
}
}
return nil
})
if transErr != nil {
return nil, transErr
}
face_detection.GlobalFaceDetector.MergeImageFaces(userOwnedImageFaceIDs, int32(destFaceGroup.ID))
return destFaceGroup, nil
}
func (r *mutationResolver) RecognizeUnlabeledFaces(ctx context.Context) ([]*models.ImageFace, error) {

View File

@@ -12,11 +12,12 @@ import (
)
type FaceDetector struct {
mutex sync.Mutex
db *gorm.DB
rec *face.Recognizer
samples []face.Descriptor
cats []int32
mutex sync.Mutex
db *gorm.DB
rec *face.Recognizer
faceDescriptors []face.Descriptor
faceGroupIDs []int32
imageFaceIDs []int
}
var GlobalFaceDetector FaceDetector
@@ -30,22 +31,23 @@ func InitializeFaceDetector(db *gorm.DB) error {
return errors.Wrap(err, "initialize facedetect recognizer")
}
samples, cats, err := getSamplesFromDatabase(db)
faceDescriptors, faceGroupIDs, imageFaceIDs, err := getSamplesFromDatabase(db)
if err != nil {
return errors.Wrap(err, "get face detection samples from database")
}
GlobalFaceDetector = FaceDetector{
db: db,
rec: rec,
samples: samples,
cats: cats,
db: db,
rec: rec,
faceDescriptors: faceDescriptors,
faceGroupIDs: faceGroupIDs,
imageFaceIDs: imageFaceIDs,
}
return nil
}
func getSamplesFromDatabase(db *gorm.DB) (samples []face.Descriptor, cats []int32, err error) {
func getSamplesFromDatabase(db *gorm.DB) (samples []face.Descriptor, faceGroupIDs []int32, imageFaceIDs []int, err error) {
var imageFaces []*models.ImageFace
@@ -54,11 +56,13 @@ func getSamplesFromDatabase(db *gorm.DB) (samples []face.Descriptor, cats []int3
}
samples = make([]face.Descriptor, len(imageFaces))
cats = make([]int32, len(imageFaces))
faceGroupIDs = make([]int32, len(imageFaces))
imageFaceIDs = make([]int, len(imageFaces))
for i, imgFace := range imageFaces {
samples[i] = face.Descriptor(imgFace.Descriptor)
cats[i] = int32(imgFace.FaceGroupID)
faceGroupIDs[i] = int32(imgFace.FaceGroupID)
imageFaceIDs[i] = imgFace.ID
}
return
@@ -150,10 +154,11 @@ func (fd *FaceDetector) classifyFace(face *face.Face, media *models.Media, image
}
}
fd.samples = append(fd.samples, face.Descriptor)
fd.cats = append(fd.cats, int32(faceGroup.ID))
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.samples, fd.cats)
fd.rec.SetSamples(fd.faceDescriptors, fd.faceGroupIDs)
return nil
}
@@ -161,19 +166,37 @@ 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
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) {
unrecognizedSamples := make([]face.Descriptor, 0)
unrecognizedCats := make([]int32, 0)
unrecognizedDescriptors := make([]face.Descriptor, 0)
unrecognizedFaceGroupIDs := make([]int32, 0)
unrecognizedImageFaceIDs := make([]int, 0)
newCats := make([]int32, 0)
newSamples := make([]face.Descriptor, 0)
newFaceGroupIDs := make([]int32, 0)
newDescriptors := make([]face.Descriptor, 0)
newImageFaceIDs := make([]int, 0)
var unlabeledFaceGroups []*models.FaceGroup
@@ -193,59 +216,64 @@ func (fd *FaceDetector) RecognizeUnlabeledFaces(tx *gorm.DB, user *models.User)
fd.mutex.Lock()
defer fd.mutex.Unlock()
for i := range fd.samples {
cat := fd.cats[i]
sample := fd.samples[i]
for i := range fd.faceDescriptors {
descriptor := fd.faceDescriptors[i]
faceGroupID := fd.faceGroupIDs[i]
imageFaceID := fd.imageFaceIDs[i]
catIsUnlabeled := false
isUnlabeled := false
for _, unlabeledFaceGroup := range unlabeledFaceGroups {
if cat == int32(unlabeledFaceGroup.ID) {
catIsUnlabeled = true
if faceGroupID == int32(unlabeledFaceGroup.ID) {
isUnlabeled = true
continue
}
}
if catIsUnlabeled {
unrecognizedCats = append(unrecognizedCats, cat)
unrecognizedSamples = append(unrecognizedSamples, sample)
if isUnlabeled {
unrecognizedFaceGroupIDs = append(unrecognizedFaceGroupIDs, faceGroupID)
unrecognizedDescriptors = append(unrecognizedDescriptors, descriptor)
unrecognizedImageFaceIDs = append(unrecognizedImageFaceIDs, imageFaceID)
} else {
newCats = append(newCats, cat)
newSamples = append(newSamples, sample)
newFaceGroupIDs = append(newFaceGroupIDs, faceGroupID)
newDescriptors = append(newDescriptors, descriptor)
newImageFaceIDs = append(newImageFaceIDs, imageFaceID)
}
}
fd.cats = newCats
fd.samples = newSamples
fd.faceGroupIDs = newFaceGroupIDs
fd.faceDescriptors = newDescriptors
fd.imageFaceIDs = newImageFaceIDs
updatedImageFaces := make([]*models.ImageFace, 0)
for i := range unrecognizedSamples {
cat := unrecognizedCats[i]
sample := unrecognizedSamples[i]
for i := range unrecognizedDescriptors {
descriptor := unrecognizedDescriptors[i]
faceGroupID := unrecognizedFaceGroupIDs[i]
imageFaceID := unrecognizedImageFaceIDs[i]
match := fd.classifyDescriptor(sample)
match := fd.classifyDescriptor(descriptor)
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)
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{
Descriptor: models.FaceDescriptor(sample),
}).First(imageFace).Error; err != nil {
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(cat)).Error; err != nil {
if err := tx.Model(&imageFace).Update("face_group_id", int(faceGroupID)).Error; err != nil {
return nil, err
}
updatedImageFaces = append(updatedImageFaces, &imageFace)
fd.cats = append(fd.cats, match)
fd.samples = append(fd.samples, sample)
fd.faceGroupIDs = append(fd.faceGroupIDs, match)
fd.faceDescriptors = append(fd.faceDescriptors, descriptor)
fd.imageFaceIDs = append(fd.imageFaceIDs, imageFaceID)
}
}

View File

@@ -4,6 +4,7 @@ import (
"log"
"net/http"
"path"
"time"
"github.com/gorilla/handlers"
"github.com/gorilla/mux"
@@ -88,6 +89,7 @@ func main() {
handler.GraphQL(photoview_graphql.NewExecutableSchema(graphqlConfig),
handler.IntrospectionEnabled(devMode),
handler.WebsocketUpgrader(server.WebsocketUpgrader(devMode)),
handler.WebsocketKeepAliveDuration(time.Second*10),
handler.WebsocketInitFunc(auth.AuthWebsocketInit(db)),
),
)