Skip to content
This repository was archived by the owner on Jan 7, 2025. It is now read-only.
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
167 changes: 110 additions & 57 deletions digits/dataset/generic/views.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,18 +7,13 @@
except ImportError:
from StringIO import StringIO

import caffe_pb2
import json
import flask
import matplotlib as mpl
import numpy as np
import PIL.Image

from .forms import GenericDatasetForm
from .job import GenericDatasetJob
from digits import extensions, utils
from digits.utils.constants import COLOR_PALETTE_ATTRIBUTE
from digits.utils.routing import request_wants_json, job_from_request
from digits.utils.lmdbreader import DbReader
from digits.utils.routing import request_wants_json, job_from_request, get_request_arg
from digits.webapp import scheduler

blueprint = flask.Blueprint(__name__, __name__)
Expand Down Expand Up @@ -135,75 +130,133 @@ def create(extension_id):
scheduler.delete_job(job)
raise

def get_database_visualizations(dataset, inputs, outputs):
# form data may be passed through the query
form_data = get_request_arg('form_data')
if form_data:
form_data = json.loads(form_data)

@blueprint.route('/explore', methods=['GET'])
# get extension ID from form and retrieve extension class
from_data = form_data and 'view_extension_id' in form_data
if from_data:
view_extension_id = form_data['view_extension_id']
else:
view_extension_id = get_request_arg('view_extension_id')

if view_extension_id:
extension_class = extensions.view.get_extension(view_extension_id)
if extension_class is None:
raise ValueError("Unknown extension '%s'" % view_extension_id)
else:
# no view extension specified, use default
extension_class = extensions.view.get_default_extension()

if from_data:
data = form_data
else:
extension_form = extension_class.get_config_form()
# validate form
extension_form_valid = extension_form.validate_on_submit()
if not extension_form_valid:
raise ValueError("Extension form validation failed with %s" % repr(extension_form.errors))
data = extension_form.data

# create instance of extension class
extension = extension_class(dataset, **data)

visualizations = []
# process data
n = len(inputs['ids'])
for idx in range(n):
input_id = inputs['ids'][idx]
input_data = inputs['data'][idx]
output_data = {key: outputs[key][idx] for key in outputs}
data = extension.process_data(
input_id,
input_data,
output_data)
template, context = extension.get_view_template(data)
visualizations.append(
flask.render_template_string(template, **context))
# get header
template, context = extension.get_header_template()
header = flask.render_template_string(template, **context) if template else None
app_begin, app_end = extension.get_ng_templates()
return visualizations, header, app_begin, app_end


@blueprint.route('/explore', methods=['GET', 'POST'])
def explore():
"""
Returns a gallery consisting of the images of one of the dbs
Returns a gallery consisting of the images from the view extension
"""
job = job_from_request()
# Get LMDB
db = job.path(flask.request.args.get('db'))
db_path = job.path(db)

if COLOR_PALETTE_ATTRIBUTE in job.extension_userdata \
and job.extension_userdata[COLOR_PALETTE_ATTRIBUTE]:
# assume single-channel 8-bit palette
palette = job.extension_userdata[COLOR_PALETTE_ATTRIBUTE]
palette = np.array(palette).reshape((len(palette) / 3, 3)) / 255.
# normalize input pixels to [0,1]
norm = mpl.colors.Normalize(vmin=0, vmax=255)
# create map
cmap = mpl.pyplot.cm.ScalarMappable(norm=norm,
cmap=mpl.colors.ListedColormap(palette))
else:
cmap = None

page = int(flask.request.args.get('page', 0))
size = int(flask.request.args.get('size', 25))

reader = DbReader(db_path)
count = 0
imgs = []
extension = extensions.data.get_extension(job.extension_id)
if extension is None:
raise ValueError("Unknown extension '%s'" % job.extension_id)

min_page = max(0, page - 5)
total_entries = reader.total_entries
# Get LMDB(s)
feature_db = job.path(flask.request.args.get('feature_db'))
feature_db_path = job.path(feature_db) if feature_db else None

label_db = job.path(flask.request.args.get('label_db'))
label_db_path = job.path(label_db) if label_db else None

# Get data from data extension
inputs, outputs, total_entries = extension.get_data(feature_db_path, label_db_path, page, size)
# Get visualization from the view extension
visualizations, header, app_begin, app_end = get_database_visualizations(job, inputs, outputs)

min_page = max(0, page - 5)
max_page = min((total_entries - 1) / size, page + 5)
pages = range(min_page, max_page + 1)
for key, value in reader.entries():
if count >= page * size:
datum = caffe_pb2.Datum()
datum.ParseFromString(value)
if not datum.encoded:
raise RuntimeError("Expected encoded database")
s = StringIO()
s.write(datum.data)
s.seek(0)
img = PIL.Image.open(s)
if cmap and img.mode in ['L', '1']:
data = np.array(img)
data = cmap.to_rgba(data) * 255
data = data.astype('uint8')
# keep RGB values only, remove alpha channel
data = data[:, :, 0:3]
img = PIL.Image.fromarray(data)
imgs.append({"label": None, "b64": utils.image.embed_image_html(img)})
count += 1
if len(imgs) >= size:
break

return flask.render_template(
'datasets/images/explore.html',
page=page, size=size, job=job, imgs=imgs, labels=None,
pages=pages, label=None, total_entries=total_entries, db=db)
# This is weak, but should do the job. This allows the form data to
# be passed to subsequent pages through the query
form_data = get_request_arg('form_data')
if not form_data:
form_data = dict(zip(flask.request.form.keys(), flask.request.form.values()))
form_data = json.dumps(form_data)

return flask.render_template('datasets/images/explore.html',
page=page, size=size, job=job, labels=None,
pages=pages, label=None, total_entries=total_entries,
feature_db=feature_db, label_db=label_db,
form_data=form_data,
header=header,
app_begin=app_begin,
app_end=app_end,
visualizations=visualizations,
)


def get_view_extensions():
"""
return all enabled view extensions
"""
view_extensions = {}
all_extensions = extensions.view.get_extensions()
for extension in all_extensions:
view_extensions[extension.get_id()] = extension.get_title()
return view_extensions


def show(job, related_jobs=None):
"""
Called from digits.dataset.views.show()
"""
return flask.render_template('datasets/generic/show.html', job=job, related_jobs=related_jobs)
data_extension = extensions.data.get_extension(job.extension_id)
can_explore = (data_extension and data_extension.can_explore())
view_extensions = get_view_extensions()
return flask.render_template('datasets/generic/show.html',
job=job,
related_jobs=related_jobs,
view_extensions=view_extensions,
can_explore=can_explore,
)


def summary(job):
Expand Down
15 changes: 8 additions & 7 deletions digits/dataset/images/classification/views.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,6 @@
from digits.utils.routing import request_wants_json, job_from_request
from digits.webapp import scheduler


blueprint = flask.Blueprint(__name__, __name__)


Expand Down Expand Up @@ -355,13 +354,15 @@ def explore():
"""
job = job_from_request()
# Get LMDB
db = flask.request.args.get('db', 'train')
db = flask.request.args.get('feature_db', 'train')
if 'train' in db.lower():
task = job.train_db_task()
elif 'val' in db.lower():
task = job.val_db_task()
elif 'test' in db.lower():
task = job.test_db_task()
else:
task = None
if task is None:
raise ValueError('No create_db task for {0}'.format(db))
if task.status != 'D':
Expand Down Expand Up @@ -416,7 +417,7 @@ def explore():
# XXX see issue #59
arr = arr[:, :, [2, 1, 0]]
img = PIL.Image.fromarray(arr)
imgs.append({"label": labels[datum.label], "b64": utils.image.embed_image_html(img)})
imgs.append({"label": labels[datum.label], "image": utils.image.embed_image_html(img)})
if label is None:
count += 1
else:
Expand All @@ -427,7 +428,7 @@ def explore():
if len(imgs) >= size:
break

return flask.render_template(
'datasets/images/explore.html',
page=page, size=size, job=job, imgs=imgs, labels=labels,
pages=pages, label=label, total_entries=total_entries, db=db)
return flask.render_template('datasets/images/explore.html',
page=page, size=size, job=job, imgs=imgs, labels=labels,
pages=pages, label=label, total_entries=total_entries,
feature_db=db, label_db=None)
74 changes: 74 additions & 0 deletions digits/extensions/data/imageSegmentation/data.py
Original file line number Diff line number Diff line change
@@ -1,16 +1,24 @@
# Copyright (c) 2016, NVIDIA CORPORATION. All rights reserved.
from __future__ import absolute_import

# Find the best implementation available
try:
from cStringIO import StringIO
except ImportError:
from StringIO import StringIO
from collections import OrderedDict
import math
import os
import random

import caffe_pb2
import numpy as np
import PIL.Image
import PIL.ImagePalette

from digits.utils import image, subclass, override, constants
from digits.utils.constants import COLOR_PALETTE_ATTRIBUTE
from digits.utils.lmdbreader import DbReader
from ..interface import DataIngestionInterface
from .forms import DatasetForm

Expand Down Expand Up @@ -226,3 +234,69 @@ def split_image_list(self, filelist, stage):
return filelist[n_val_entries:]
else:
raise ValueError("Unknown stage: %s" % stage)

@staticmethod
@override
def can_explore():
return True

@staticmethod
def get_data_from_db(db_path, page, size):
reader = DbReader(db_path)
data = []
count = 0
for key, value in reader.entries():
if count >= page * size:
datum = caffe_pb2.Datum()
datum.ParseFromString(value)
if not datum.encoded:
raise RuntimeError("Expected encoded database")
s = StringIO()
s.write(datum.data)
s.seek(0)
datum = PIL.Image.open(s)
datum = np.array(datum)
data.append(datum)
count += 1
if len(data) >= size:
break
return data, reader.total_entries

@staticmethod
def convert(data):
"""
Convert the labael data to look like inference data
:param data: label data
:return: output
- output is the data converted to look like inference data
"""
idxs = np.unique(data)
mx = np.max(idxs[np.where(idxs != 255)])
output = [np.equal(data, i).astype('float') for i in range(mx + 1)]
output = np.array(output)
return output

@staticmethod
@override
def get_data(feature_db_path, label_db_path, page, size):
"""
Retrieve data from databases and return it as if it were inference data
so that view extensions can present it in the Explore DB pages.
:param feature_db_path: feature database path
:param label_db_path: label database path
:param page: which page
:param size: number of items per page
:return: (inputs, outputs, total_entries)
- inputs are the feature data
- outputs are the label data made to look like inference data
- total_entries is the number of entries in the feature database
"""
features, total_entries = DataIngestion.get_data_from_db(feature_db_path, page, size)
inputs = {'data': features, 'ids': range(len(features))}

labels, _ = DataIngestion.get_data_from_db(label_db_path, page, size)
labels = [DataIngestion.convert(label) for label in labels]
labels = np.array(labels)
outputs = OrderedDict([('score', labels)])

return inputs, outputs, total_entries
24 changes: 24 additions & 0 deletions digits/extensions/data/interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -110,3 +110,27 @@ def itemize_entries(self, stage):
this function in no particular order
"""
raise NotImplementedError

@staticmethod
def can_explore():

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

why don't you have get_data() return None by default and check the return value on caller side? This way, sub-classes only have to override get_data()

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I didn't because get_data is expensive to compute and the data would be thrown away in that context.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Also, when all the view extensions are converted, then the can_explore might just go away.

"""
To be overridden
:return: True if explore methods are provided
"""
return False

@staticmethod
def get_data(feature_db_path, label_db_path, page, size):
"""
Retrieves a page worth of data from the databases and returns
data to be passed to a view extension
:param feature_db_path: full path to features database
:param label_db_path: full path to labels database
:param page: page number requested
:param size: number of items per page
:return: (inputs, outputs, total_entries)
- inputs: features
- output: labels
- total_entries: number of entries in the database
"""
raise NotImplementedError
Loading