mirror of
https://git.vectorsigma.ru/public/photoview.git
synced 2026-08-03 20:59:03 +00:00
Add UI to move faces between face groups
This commit is contained in:
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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)),
|
||||
),
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user