diff --git a/ComixGAN/G_weights.best.hdf5 b/ComixGAN/G_weights.best.hdf5 new file mode 100644 index 0000000..e73777f Binary files /dev/null and b/ComixGAN/G_weights.best.hdf5 differ diff --git a/ComixGAN/model.py b/ComixGAN/model.py new file mode 100644 index 0000000..7cca114 --- /dev/null +++ b/ComixGAN/model.py @@ -0,0 +1,61 @@ +from keras import layers, Model +from keras_contrib.layers import InstanceNormalization + + +class ComixGAN(): + def __init__(self): + # Build and compile the generator + self.generator = self.build_generator() + + def build_generator(self): + def residual_block(input_tensor): + c1 = layers.Conv2D(256, (3, 3), strides=(1, 1), padding='same', use_bias=False)(input_tensor) + bn1 = InstanceNormalization(axis=3)(c1) + a1 = layers.Activation('relu')(bn1) + c2 = layers.Conv2D(256, (3, 3), strides=(1, 1), padding='same', use_bias=False)(a1) + bn2 = InstanceNormalization(axis=3)(c2) + add1 = layers.add([bn2, input_tensor]) + + return add1 + + inp = layers.Input((None, None, 3)) + + c1 = layers.Conv2D(64, (7, 7), strides=(1, 1), padding='same', use_bias=False)(inp) + bn1 = InstanceNormalization(axis=3)(c1) + a1 = layers.Activation('relu')(bn1) + + c2 = layers.Conv2D(128, (3, 3), strides=(2, 2), padding='same')(a1) + c3 = layers.Conv2D(128, (3, 3), strides=(1, 1), padding='same', use_bias=False)(c2) + bn2 = InstanceNormalization(axis=3)(c3) + a2 = layers.Activation('relu')(bn2) + + c4 = layers.Conv2D(256, (3, 3), strides=(2, 2), padding='same')(a2) + c5 = layers.Conv2D(256, (3, 3), strides=(1, 1), padding='same', use_bias=False)(c4) + bn3 = InstanceNormalization(axis=3)(c5) + a3 = layers.Activation('relu')(bn3) + + r1 = residual_block(a3) + r2 = residual_block(r1) + r3 = residual_block(r2) + r4 = residual_block(r3) + r5 = residual_block(r4) + r6 = residual_block(r5) + r7 = residual_block(r6) + r8 = residual_block(r7) + + u1 = layers.UpSampling2D(size=(2, 2))(r8) + c6 = layers.Conv2D(128, (3, 3), strides=(1, 1), padding='same', use_bias=False)(u1) + bn4 = InstanceNormalization(axis=3)(c6) + a4 = layers.Activation('relu')(bn4) + + u2 = layers.UpSampling2D(size=(2, 2))(a4) + c7 = layers.Conv2D(64, (3, 3), strides=(1, 1), padding='same', use_bias=False)(u2) + bn5 = InstanceNormalization(axis=3)(c7) + a5 = layers.Activation('relu')(bn5) + + output = layers.Conv2D(3, (7, 7), strides=(1, 1), activation='sigmoid', padding='same')(a5) + + return Model(inputs=[inp], outputs=[output]) + + +comixGAN = ComixGAN() diff --git a/api/models.py b/api/models.py index ee81408..9b4a18f 100644 --- a/api/models.py +++ b/api/models.py @@ -34,14 +34,15 @@ class Video(models.Model): else: raise TooLargeFile() - def create_comic(self, frames_mode=0, rl_mode=0, image_assessment_mode=0): + def create_comic(self, frames_mode=0, rl_mode=0, image_assessment_mode=0, style_transfer_mode=0): (keyframes, keyframes_timings), keyframes_extraction_time = KeyFramesExtractor.get_keyframes( video=self, frames_mode=frames_mode, rl_mode=rl_mode, image_assessment_mode=image_assessment_mode ) - stylized_keyframes, stylization_time = StyleTransfer.get_stylized_frames(frames=keyframes) + stylized_keyframes, stylization_time = StyleTransfer.get_stylized_frames(frames=keyframes, + style_transfer_mode=style_transfer_mode) comic_image, layout_generation_time = LayoutGenerator.get_layout(frames=stylized_keyframes) timings = { @@ -50,6 +51,7 @@ class Video(models.Model): 'layout_generation_time': layout_generation_time, **keyframes_timings } + return comic_image, timings @@ -61,7 +63,7 @@ class Comic(models.Model): @profile def create_from_nparray(cls, nparray_file, video): if nparray_file.max() <= 1: - nparray_file = (nparray_file * 255).astype(int) + nparray_file = (nparray_file).astype(int) tmp_name = uuid.uuid4().hex + ".png" cv2.imwrite(jj(settings.TMP_DIR, tmp_name), nparray_file) with open(jj(settings.TMP_DIR, tmp_name), mode="rb") as tmp_file: diff --git a/api/serializers.py b/api/serializers.py index 47c7c05..5510527 100644 --- a/api/serializers.py +++ b/api/serializers.py @@ -8,7 +8,10 @@ class VideoSerializer(serializers.Serializer): file = serializers.FileField() frames_mode = serializers.IntegerField(min_value=0, max_value=1, default=settings.DEFAULT_FRAMES_SAMPLING_MODE) rl_mode = serializers.IntegerField(min_value=0, max_value=1, default=settings.DEFAULT_RL_MODE) - image_assessment_mode = serializers.IntegerField(min_value=0, max_value=1, default=settings.DEFAULT_IMAGE_ASSESSMENT_MODE) + image_assessment_mode = serializers.IntegerField(min_value=0, max_value=1, + default=settings.DEFAULT_IMAGE_ASSESSMENT_MODE) + style_transfer_mode = serializers.IntegerField(min_value=0, max_value=2, + default=settings.DEFAULT_STYLE_TRANSFER_MODE) def validate(self, attrs): file = attrs.get("file") @@ -23,4 +26,7 @@ class YouTubeDownloadSerializer(serializers.Serializer): url = serializers.URLField() frames_mode = serializers.IntegerField(min_value=0, max_value=1, default=settings.DEFAULT_FRAMES_SAMPLING_MODE) rl_mode = serializers.IntegerField(min_value=0, max_value=1, default=settings.DEFAULT_RL_MODE) - image_assessment_mode = serializers.IntegerField(min_value=0, max_value=1, default=settings.DEFAULT_IMAGE_ASSESSMENT_MODE) + image_assessment_mode = serializers.IntegerField(min_value=0, max_value=1, + default=settings.DEFAULT_IMAGE_ASSESSMENT_MODE) + style_transfer_mode = serializers.IntegerField(min_value=0, max_value=2, + default=settings.DEFAULT_STYLE_TRANSFER_MODE) diff --git a/api/views.py b/api/views.py index 77cde9d..add133c 100644 --- a/api/views.py +++ b/api/views.py @@ -22,7 +22,8 @@ class Comixify(APIView): comic_image, timings = video.create_comic( frames_mode=serializer.validated_data["frames_mode"], rl_mode=serializer.validated_data["rl_mode"], - image_assessment_mode=serializer.validated_data["image_assessment_mode"] + image_assessment_mode=serializer.validated_data["image_assessment_mode"], + style_transfer_mode=serializer.validated_data["rl_mode"], ) comic, from_nparray_time = Comic.create_from_nparray(comic_image, video) timings['from_nparray_time'] = from_nparray_time @@ -54,7 +55,8 @@ class ComixifyFromYoutube(APIView): comic_image, timings = video.create_comic( frames_mode=serializer.validated_data["frames_mode"], rl_mode=serializer.validated_data["rl_mode"], - image_assessment_mode=serializer.validated_data["image_assessment_mode"] + image_assessment_mode=serializer.validated_data["image_assessment_mode"], + style_transfer_mode=serializer.validated_data["rl_mode"], ) comic, from_nparray_time = Comic.create_from_nparray(comic_image, video) timings['from_nparray_time'] = from_nparray_time diff --git a/comic_layout/comic_layout.py b/comic_layout/comic_layout.py index 27a9dad..07a0f9a 100644 --- a/comic_layout/comic_layout.py +++ b/comic_layout/comic_layout.py @@ -28,6 +28,6 @@ class LayoutGenerator(): def _pad_images(frames): padded_result_imgs = [] for img in frames: - padded_img = cv2.copyMakeBorder(img, 5, 5, 5, 5, cv2.BORDER_CONSTANT, value=(1, 1, 1)) + padded_img = cv2.copyMakeBorder(img, 5, 5, 5, 5, cv2.BORDER_CONSTANT, value=(255, 255, 255)) padded_result_imgs.append(padded_img) return padded_result_imgs diff --git a/dockerfile b/dockerfile index 718986e..48429f1 100644 --- a/dockerfile +++ b/dockerfile @@ -10,10 +10,7 @@ RUN apt-get update && apt-get install -y apt-utils software-properties-common && libsnappy-dev protobuf-compiler \ python-numpy python-setuptools python-scipy \ libavformat-dev libswscale-dev unzip && \ - python3.6 -m pip install --upgrade pip && \ - python3.6 -m pip install jupyter ipywidgets jupyterlab && \ - python3.6 -m pip install h5py keras && \ - python3.6 -m pip install scikit-image opencv-contrib-python pyyaml + python3.6 -m pip install --upgrade pip RUN mkdir /comixify COPY ./Makefile.config /comixify/Makefile.config @@ -46,6 +43,7 @@ WORKDIR /comixify COPY . /comixify RUN unzip popularity/pretrained_model/svr_test_11.10.sk.zip -d popularity/pretrained_model/ && \ python3.6 -m pip install -r requirements.txt + python3.6 -m pip install git+https://www.github.com/keras-team/keras-contrib.git # Port to expose diff --git a/frontend/client/App.js b/frontend/client/App.js index db415a3..016c403 100644 --- a/frontend/client/App.js +++ b/frontend/client/App.js @@ -32,7 +32,8 @@ class App extends React.Component { result_comics: null, framesMode: "0", rlMode: "0", - imageAssessment: "0" + imageAssessment: "0", + styleTransferMode: "0", }; this.onVideoDrop = this.onVideoDrop.bind(this); this.onModelChange = this.onModelChange.bind(this); @@ -40,6 +41,7 @@ class App extends React.Component { this.onYouTubeSubmit = this.onYouTubeSubmit.bind(this); this.onSamplingChange = this.onSamplingChange.bind(this); this.onImageAssessmentChange = this.onImageAssessmentChange.bind(this); + this.styleTransferChange = this.styleTransferChange.bind(this); } static onVideoUploadProgress(progressEvent) { let percentCompleted = Math.round( @@ -52,6 +54,12 @@ class App extends React.Component { this.setState({ rlMode: value }) + } + styleTransferChange(e) { + let value = e.currentTarget.value; + this.setState({ + styleTransferMode: value + }) } onSamplingChange(e) { let value = e.currentTarget.value; @@ -78,12 +86,13 @@ class App extends React.Component { } } processVideo(video) { - let { framesMode, rlMode, imageAssessment } = this.state; + let { framesMode, rlMode, imageAssessment, styleTransferMode } = this.state; let data = new FormData(); data.append("file", video); data.set('frames_mode', parseInt(framesMode)); data.set('rl_mode', parseInt(rlMode)); data.set("image_assessment_mode", parseInt(imageAssessment)); + data.set('style_transfer_mode', parseInt(styleTransferMode)); post(COMIXIFY_API, data, { headers: { "content-type": "multipart/form-data" }, onUploadProgress: App.onVideoUploadProgress @@ -111,12 +120,13 @@ class App extends React.Component { this.processVideo(files[0]); } submitYouTube(link) { - let { framesMode, rlMode, imageAssessment } = this.state; + let { framesMode, rlMode, imageAssessment, styleTransferMode } = this.state; post(FROM_YOUTUBE_API, { url: link, frames_mode: parseInt(framesMode), rl_mode: parseInt(rlMode), - image_assessment_mode: parseInt(imageAssessment) + image_assessment_mode: parseInt(imageAssessment), + style_transfer_mode: parseInt(styleTransferMode) }) .then(this.handleResponse) .catch(err => { @@ -143,7 +153,7 @@ class App extends React.Component { } render() { let { - state, drop_errors, result_comics, framesMode, rlMode, videoId, imageAssessment + state, drop_errors, result_comics, framesMode, rlMode, videoId, imageAssessment, styleTransferMode } = this.state; let showUsage = [ App.appStates.INITIAL, @@ -231,6 +241,36 @@ class App extends React.Component { onChange={this.onImageAssessmentChange} /> + +
+ Style Transfer model: + + + + + +
)} diff --git a/frontend/templates/index.html b/frontend/templates/index.html index 7672b44..c1225e4 100644 --- a/frontend/templates/index.html +++ b/frontend/templates/index.html @@ -60,7 +60,7 @@ * Max size of video is 50 MB -

Respose:

+

Response:

{
diff --git a/requirements.txt b/requirements.txt index 7d4a461..7f4cfc2 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,16 +1,32 @@ +absl-py==0.6.1 +astor==0.7.1 Django==2.0.7 django-rest-framework==0.1.0 djangorestframework==3.8.2 +gast==0.2.0 +grpcio==1.16.0 gunicorn==19.9.0 +h5py==2.8.0 +Keras==2.2.2 +Keras-Applications==1.0.6 +keras-contrib==2.0.8 +Keras-Preprocessing==1.0.5 +Markdown==3.0.1 numpy==1.14.5 opencv-python==3.4.2.17 pafy==0.5.4 Pillow==5.2.0 +protobuf==3.6.1 psycopg2==2.7.5 pytz==2018.5 +PyYAML==3.13 +scipy==1.1.0 six==1.11.0 +tensorboard==1.11.0 +tensorflow-gpu==1.11.0 +termcolor==1.1.0 torch==0.4.1 torchvision==0.2.1 -scikit-learn==0.19.2 +Werkzeug==0.14.1 youtube-dl==2018.9.18 tensorflow-gpu==1.10.1 \ No newline at end of file diff --git a/settings/settings.py b/settings/settings.py index 5f90b93..658ff10 100644 --- a/settings/settings.py +++ b/settings/settings.py @@ -1,148 +1,150 @@ -""" -Django settings for comixify project. - -Generated by 'django-admin startproject' using Django 2.0.7. - -For more information on this file, see -https://docs.djangoproject.com/en/2.0/topics/settings/ - -For the full list of settings and their values, see -https://docs.djangoproject.com/en/2.0/ref/settings/ -""" - -import os - -# Build paths inside the project like this: os.path.join(BASE_DIR, ...) -BASE_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - -# Quick-start development settings - unsuitable for production -# See https://docs.djangoproject.com/en/2.0/howto/deployment/checklist/ - -# SECURITY WARNING: keep the secret key used in production secret! -# Use a separate file for the secret key -with open(os.path.join(BASE_DIR, 'secretkey.txt')) as f: - SECRET_KEY = f.read().strip() - -# SECURITY WARNING: don't run with debug turned on in production! -DEBUG = os.environ.get('DEBUG') == 'true' - -ALLOWED_HOSTS = ['35.241.250.34', 'comixify.ii.pw.edu.pl', 'localhost', '127.0.0.1'] - -# Application definition - -INSTALLED_APPS = [ - 'django.contrib.admin', - 'django.contrib.auth', - 'django.contrib.contenttypes', - 'django.contrib.sessions', - 'django.contrib.messages', - 'django.contrib.staticfiles', - 'rest_framework', - 'api', - 'style_transfer', - 'comic_layout', - 'frontend', -] - -MIDDLEWARE = [ - 'django.middleware.security.SecurityMiddleware', - 'django.contrib.sessions.middleware.SessionMiddleware', - 'django.middleware.common.CommonMiddleware', - 'django.middleware.csrf.CsrfViewMiddleware', - 'django.contrib.auth.middleware.AuthenticationMiddleware', - 'django.contrib.messages.middleware.MessageMiddleware', - 'django.middleware.clickjacking.XFrameOptionsMiddleware', -] - -ROOT_URLCONF = 'settings.urls' - -TEMPLATES = [ - { - 'BACKEND': 'django.template.backends.django.DjangoTemplates', - 'DIRS': [], - 'APP_DIRS': True, - 'OPTIONS': { - 'context_processors': [ - 'django.template.context_processors.debug', - 'django.template.context_processors.request', - 'django.contrib.auth.context_processors.auth', - 'django.contrib.messages.context_processors.messages', - ], - }, - }, -] - -WSGI_APPLICATION = 'settings.wsgi.application' - -# Database -# https://docs.djangoproject.com/en/2.0/ref/settings/#databases - -DATABASES = { - 'default': { - 'ENGINE': 'django.db.backends.postgresql_psycopg2', - 'NAME': 'postgres', - 'USER': 'postgres', - 'HOST': 'db', - 'PORT': '5432', - 'PASSWORD': 'postgres', - 'CONN_MAX_AGE': 60, - } -} - -# Password validation -# https://docs.djangoproject.com/en/2.0/ref/settings/#auth-password-validators - -AUTH_PASSWORD_VALIDATORS = [ - { - 'NAME': 'django.contrib.auth.password_validation.UserAttributeSimilarityValidator', - }, - { - 'NAME': 'django.contrib.auth.password_validation.MinimumLengthValidator', - }, - { - 'NAME': 'django.contrib.auth.password_validation.CommonPasswordValidator', - }, - { - 'NAME': 'django.contrib.auth.password_validation.NumericPasswordValidator', - }, -] -CACHES = { - 'default': { - 'BACKEND': 'django.core.cache.backends.locmem.LocMemCache', - 'LOCATION': 'unique-snowflake', - } -} - -# Internationalization -# https://docs.djangoproject.com/en/2.0/topics/i18n/ - -LANGUAGE_CODE = 'en-us' - -TIME_ZONE = 'UTC' - -USE_I18N = True - -USE_L10N = True - -USE_TZ = True - -# Static files (CSS, JavaScript, Images) -# https://docs.djangoproject.com/en/2.0/howto/static-files/ - -STATIC_URL = '/static/' -STATIC_ROOT = os.path.join(BASE_DIR, 'static') - -MEDIA_URL = '/media/' -MEDIA_ROOT = os.path.join(BASE_DIR, 'media') - -PERMITTED_VIDEO_EXTENSIONS = ['mp4', 'avi'] -MAX_FILE_SIZE = 50000000 -NUMBERS_OF_FRAMES_TO_SHOW = 10 -TMP_DIR = 'tmp/' -GPU = True - -FEATURE_BATCH_SIZE = 32 -DEFAULT_FRAMES_SAMPLING_MODE = 0 -DEFAULT_RL_MODE = 0 -DEFAULT_IMAGE_ASSESSMENT_MODE = 0 -DEBUG = True +""" +Django settings for comixify project. + +Generated by 'django-admin startproject' using Django 2.0.7. + +For more information on this file, see +https://docs.djangoproject.com/en/2.0/topics/settings/ + +For the full list of settings and their values, see +https://docs.djangoproject.com/en/2.0/ref/settings/ +""" + +import os + +# Build paths inside the project like this: os.path.join(BASE_DIR, ...) +BASE_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + +# Quick-start development settings - unsuitable for production +# See https://docs.djangoproject.com/en/2.0/howto/deployment/checklist/ + +# SECURITY WARNING: keep the secret key used in production secret! +# Use a separate file for the secret key +with open(os.path.join(BASE_DIR, 'secretkey.txt')) as f: + SECRET_KEY = f.read().strip() + +# SECURITY WARNING: don't run with debug turned on in production! +DEBUG = os.environ.get('DEBUG') == 'true' + +ALLOWED_HOSTS = ['35.241.250.34', 'comixify.ii.pw.edu.pl', 'localhost', '127.0.0.1'] + +# Application definition + +INSTALLED_APPS = [ + 'django.contrib.admin', + 'django.contrib.auth', + 'django.contrib.contenttypes', + 'django.contrib.sessions', + 'django.contrib.messages', + 'django.contrib.staticfiles', + 'rest_framework', + 'api', + 'style_transfer', + 'comic_layout', + 'frontend', +] + +MIDDLEWARE = [ + 'django.middleware.security.SecurityMiddleware', + 'django.contrib.sessions.middleware.SessionMiddleware', + 'django.middleware.common.CommonMiddleware', + 'django.middleware.csrf.CsrfViewMiddleware', + 'django.contrib.auth.middleware.AuthenticationMiddleware', + 'django.contrib.messages.middleware.MessageMiddleware', + 'django.middleware.clickjacking.XFrameOptionsMiddleware', +] + +ROOT_URLCONF = 'settings.urls' + +TEMPLATES = [ + { + 'BACKEND': 'django.template.backends.django.DjangoTemplates', + 'DIRS': [], + 'APP_DIRS': True, + 'OPTIONS': { + 'context_processors': [ + 'django.template.context_processors.debug', + 'django.template.context_processors.request', + 'django.contrib.auth.context_processors.auth', + 'django.contrib.messages.context_processors.messages', + ], + }, + }, +] + +WSGI_APPLICATION = 'settings.wsgi.application' + +# Database +# https://docs.djangoproject.com/en/2.0/ref/settings/#databases + +DATABASES = { + 'default': { + 'ENGINE': 'django.db.backends.postgresql_psycopg2', + 'NAME': 'postgres', + 'USER': 'postgres', + 'HOST': 'db', + 'PORT': '5432', + 'PASSWORD': 'postgres', + 'CONN_MAX_AGE': 60, + } +} + +# Password validation +# https://docs.djangoproject.com/en/2.0/ref/settings/#auth-password-validators + +AUTH_PASSWORD_VALIDATORS = [ + { + 'NAME': 'django.contrib.auth.password_validation.UserAttributeSimilarityValidator', + }, + { + 'NAME': 'django.contrib.auth.password_validation.MinimumLengthValidator', + }, + { + 'NAME': 'django.contrib.auth.password_validation.CommonPasswordValidator', + }, + { + 'NAME': 'django.contrib.auth.password_validation.NumericPasswordValidator', + }, +] +CACHES = { + 'default': { + 'BACKEND': 'django.core.cache.backends.locmem.LocMemCache', + 'LOCATION': 'unique-snowflake', + } +} + +# Internationalization +# https://docs.djangoproject.com/en/2.0/topics/i18n/ + +LANGUAGE_CODE = 'en-us' + +TIME_ZONE = 'UTC' + +USE_I18N = True + +USE_L10N = True + +USE_TZ = True + +# Static files (CSS, JavaScript, Images) +# https://docs.djangoproject.com/en/2.0/howto/static-files/ + +STATIC_URL = '/static/' +STATIC_ROOT = os.path.join(BASE_DIR, 'static') + +MEDIA_URL = '/media/' +MEDIA_ROOT = os.path.join(BASE_DIR, 'media') + +PERMITTED_VIDEO_EXTENSIONS = ['mp4', 'avi'] +MAX_FILE_SIZE = 50000000 +NUMBERS_OF_FRAMES_TO_SHOW = 10 +TMP_DIR = 'tmp/' +GPU = True + +FEATURE_BATCH_SIZE = 32 +DEFAULT_FRAMES_SAMPLING_MODE = 0 +DEFAULT_RL_MODE = 0 +DEFAULT_IMAGE_ASSESSMENT_MODE = 0 + +DEFAULT_STYLE_TRANSFER_MODE = 0 +COMIX_GAN_WEIGHTS_PATH = os.path.join(BASE_DIR, 'ComixGAN', 'G_weights.best.hdf5') diff --git a/style_transfer/style_transfer.py b/style_transfer/style_transfer.py index e4b5c45..3e1135a 100644 --- a/style_transfer/style_transfer.py +++ b/style_transfer/style_transfer.py @@ -10,20 +10,56 @@ from torch.autograd import Variable from CartoonGAN.network.Transformer import Transformer from utils import profile +from ComixGAN.model import ComixGAN class StyleTransfer(): @classmethod @profile - def get_stylized_frames(cls, frames, method="cartoon_gan", gpu=settings.GPU, **kwargs): - if method == "cartoon_gan": - return cls._cartoon_gan_stylize(frames, gpu=gpu, **kwargs) + def get_stylized_frames(cls, frames, style_transfer_mode=0, gpu=settings.GPU): + if style_transfer_mode == 0: + return cls._comix_gan_stylize(frames=frames) + elif style_transfer_mode == 1: + return cls._cartoon_gan_stylize(frames, gpu=gpu, style='Hayao') + elif style_transfer_mode == 2: + return cls._cartoon_gan_stylize(frames, gpu=gpu, style='Hosoda') @staticmethod - def _cartoon_gan_stylize(frames, gpu=True, **kwargs): - style = kwargs.get("style", "Hayao") - resize = kwargs.get("resize", 450) + def _resize_images(frames, size=384): + resized_images = [] + for img in frames: + # resize image, keep aspect ratio + h, w, _ = img.shape + ratio = h / w + if ratio > 1: + h = size + w = int(h * 1.0 / ratio) + else: + w = size + h = int(w * ratio) + resized_img = cv2.resize(img, (w, h)) + resized_images.append(resized_img) + return resized_images + @classmethod + def _comix_gan_stylize(cls, frames): + model_cache_key = 'model_cache' + model = cache.get(model_cache_key) # get model from cache + + if model is None: + # load pretrained model + model = ComixGAN() + model.load_weights(settings.COMIX_GAN_WEIGHTS_PATH) + cache.set(model_cache_key, model, None) # None is the timeout parameter. It means cache forever + + frames = cls._resize_images(frames, size=384) + frames = np.stack(frames) + frames = frames / frames.max() + stylized_imgs = model.generator.predict(frames) + return list(255 * stylized_imgs) + + @classmethod + def _cartoon_gan_stylize(cls, frames, gpu=True, style='Hayao'): model_cache_key = 'model_cache' model = cache.get(model_cache_key) # get model from cache @@ -35,19 +71,10 @@ class StyleTransfer(): model.cuda() if gpu else model.float() cache.set(model_cache_key, model, None) # None is the timeout parameter. It means cache forever + frames = cls._resize_images(frames, size=450) stylized_imgs = [] for img in frames: - # resize image, keep aspect ratio - h, w, _ = img.shape - ratio = h * 1.0 / w - if ratio > 1: - h = resize - w = int(h * 1.0 / ratio) - else: - w = resize - h = int(w * ratio) - input_image = cv2.resize(img, (w, h)) - input_image = transforms.ToTensor()(input_image).unsqueeze(0) + input_image = transforms.ToTensor()(img).unsqueeze(0) # preprocess, (-1, 1) input_image = -1 + 2 * input_image @@ -64,6 +91,6 @@ class StyleTransfer(): output_image = np.rollaxis(output_image, 0, 3) # append image to result images - stylized_imgs.append(output_image) + stylized_imgs.append(255 * output_image) return stylized_imgs