Skip to content
Merged
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
81 changes: 44 additions & 37 deletions docs/conf.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,11 +18,13 @@
# documentation root, use os.path.abspath to make it absolute, like shown here.
#
import os
import pkg_resources
import sys
sys.path.insert(0, os.path.abspath('..'))

__version__ = pkg_resources.get_distribution('pfrl').version
import pkg_resources

sys.path.insert(0, os.path.abspath(".."))

__version__ = pkg_resources.get_distribution("pfrl").version


# -- General configuration ------------------------------------------------
Expand All @@ -34,31 +36,33 @@
# Add any Sphinx extension module names here, as strings. They can be
# extensions coming with Sphinx (named 'sphinx.ext.*') or your custom
# ones.
extensions = ['sphinx.ext.autodoc',
'sphinx.ext.doctest',
'sphinx.ext.intersphinx',
'sphinx.ext.todo',
'sphinx.ext.coverage',
'sphinx.ext.mathjax',
'sphinx.ext.napoleon',
'sphinx.ext.viewcode']
extensions = [
"sphinx.ext.autodoc",
"sphinx.ext.doctest",
"sphinx.ext.intersphinx",
"sphinx.ext.todo",
"sphinx.ext.coverage",
"sphinx.ext.mathjax",
"sphinx.ext.napoleon",
"sphinx.ext.viewcode",
]

# Add any paths that contain templates here, relative to this directory.
templates_path = ['_templates']
templates_path = ["_templates"]

# The suffix(es) of source filenames.
# You can specify multiple suffix as a list of string:
#
# source_suffix = ['.rst', '.md']
source_suffix = '.rst'
source_suffix = ".rst"

# The master toctree document.
master_doc = 'index'
master_doc = "index"

# General information about the project.
project = 'PFRL'
copyright = '2020, Preferred Networks, Inc.'
author = 'Preferred Networks, Inc.'
project = "PFRL"
copyright = "2020, Preferred Networks, Inc."
author = "Preferred Networks, Inc."

# The version info for the project you're documenting, acts as replacement for
# |version| and |release|, also used in various other places throughout the
Expand All @@ -79,10 +83,10 @@
# List of patterns, relative to source directory, that match files and
# directories to ignore when looking for source files.
# This patterns also effect to html_static_path and html_extra_path
exclude_patterns = ['_build', 'Thumbs.db', '.DS_Store']
exclude_patterns = ["_build", "Thumbs.db", ".DS_Store"]

# The name of the Pygments (syntax highlighting) style to use.
pygments_style = 'sphinx'
pygments_style = "sphinx"

# If true, `todo` and `todoList` produce output, else they produce nothing.
todo_include_todos = True
Expand All @@ -93,7 +97,7 @@
# The theme to use for HTML and HTML Help pages. See the documentation for
# a list of builtin themes.
#
html_theme = 'sphinx_rtd_theme'
html_theme = "sphinx_rtd_theme"

# Theme options are theme-specific and customize the look and feel of a theme
# further. For a list of options available for each theme, see the
Expand All @@ -104,13 +108,13 @@
# Add any paths that contain custom static files (such as style sheets) here,
# relative to this directory. They are copied after the builtin static files,
# so a file named "default.css" will overwrite the builtin "default.css".
html_static_path = ['_static']
html_static_path = ["_static"]


# -- Options for HTMLHelp output ------------------------------------------

# Output file base name for HTML help builder.
htmlhelp_basename = 'PFRLdoc'
htmlhelp_basename = "PFRLdoc"


# -- Options for LaTeX output ---------------------------------------------
Expand All @@ -119,15 +123,12 @@
# The paper size ('letterpaper' or 'a4paper').
#
# 'papersize': 'letterpaper',

# The font size ('10pt', '11pt' or '12pt').
#
# 'pointsize': '10pt',

# Additional stuff for the LaTeX preamble.
#
# 'preamble': '',

# Latex figure (float) alignment
#
# 'figure_align': 'htbp',
Expand All @@ -137,19 +138,21 @@
# (source start file, target name, title,
# author, documentclass [howto, manual, or own class]).
latex_documents = [
(master_doc, 'PFRL.tex', 'PFRL Documentation',
'Preferred Networks, Inc.', 'manual'),
(
master_doc,
"PFRL.tex",
"PFRL Documentation",
"Preferred Networks, Inc.",
"manual",
),
]


# -- Options for manual page output ---------------------------------------

# One entry per manual page. List of tuples
# (source start file, name, description, authors, manual section).
man_pages = [
(master_doc, 'pfrl', 'PFRL Documentation',
[author], 1)
]
man_pages = [(master_doc, "pfrl", "PFRL Documentation", [author], 1)]


# -- Options for Texinfo output -------------------------------------------
Expand All @@ -158,13 +161,17 @@
# (source start file, target name, title, author,
# dir menu entry, description, category)
texinfo_documents = [
(master_doc, 'PFRL', 'PFRL Documentation',
author, 'PFRL', 'One line description of project.',
'Miscellaneous'),
(
master_doc,
"PFRL",
"PFRL Documentation",
author,
"PFRL",
"One line description of project.",
"Miscellaneous",
),
]




# Example configuration for intersphinx: refer to the Python standard library.
intersphinx_mapping = {'https://docs.python.org/': None}
intersphinx_mapping = {"https://docs.python.org/": None}
1 change: 1 addition & 0 deletions examples/atari/train_drqn_ale.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
To train DQRN using a recurrent model on flickering 1-frame Breakout, run:
python train_drqn_ale.py --recurrent --flicker --no-frame-stack
"""

import argparse

import gym
Expand Down
1 change: 1 addition & 0 deletions examples/atari/train_ppo_ale.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
To train PPO using a recurrent model on a flickering Atari env, run:
python train_ppo_ale.py --recurrent --flicker --no-frame-stack
"""

import argparse
import functools

Expand Down
1 change: 1 addition & 0 deletions examples/atlas/train_soft_actor_critic_atlas.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
"""A training script of Soft Actor-Critic on RoboschoolAtlasForwardWalk-v1."""

import argparse
import functools
import logging
Expand Down
1 change: 1 addition & 0 deletions examples/gym/train_reinforce_gym.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
To solve InvertedPendulum-v1, run:
python train_reinforce_gym.py --env InvertedPendulum-v1
"""

import argparse

import gym
Expand Down
1 change: 1 addition & 0 deletions examples/mujoco/reproduction/ppo/train_ppo.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
This script follows the settings of https://arxiv.org/abs/1709.06560 as much
as possible.
"""

import argparse
import functools

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
This script follows the settings of https://arxiv.org/abs/1812.05905 as much
as possible.
"""

import argparse
import functools
import logging
Expand Down
1 change: 1 addition & 0 deletions examples/mujoco/reproduction/trpo/train_trpo.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
This script follows the settings of https://arxiv.org/abs/1709.06560 as much
as possible.
"""

import argparse
import logging

Expand Down
4 changes: 0 additions & 4 deletions pfrl/agents/iqn.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,6 @@ def cosine_basis_functions(x, n_basis_functions=64):


class CosineBasisLinear(nn.Module):

"""Linear layer following cosine basis functions.

Args:
Expand Down Expand Up @@ -81,7 +80,6 @@ def _evaluate_psi_x_with_quantile_thresholds(psi_x, phi, f, taus):


class ImplicitQuantileQFunction(nn.Module):

"""Implicit quantile network-based Q-function.

Args:
Expand Down Expand Up @@ -125,7 +123,6 @@ def evaluate_with_quantile_thresholds(taus):


class RecurrentImplicitQuantileQFunction(Recurrent, nn.Module):

"""Recurrent implicit quantile network-based Q-function.

Args:
Expand Down Expand Up @@ -256,7 +253,6 @@ def compute_weighted_value_loss(eltwise_loss, weights, batch_accumulator="mean")


class IQN(dqn.DQN):

"""Implicit Quantile Networks.

See https://arxiv.org/abs/1806.06923.
Expand Down
4 changes: 2 additions & 2 deletions pfrl/initializers/chainer_default.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
"""Initializes the weights and biases of a layer to chainer default.
"""
"""Initializes the weights and biases of a layer to chainer default."""

import torch
import torch.nn as nn

Expand Down
46 changes: 24 additions & 22 deletions setup.py
Original file line number Diff line number Diff line change
@@ -1,30 +1,32 @@
import codecs
from setuptools import find_packages
from setuptools import setup

from setuptools import find_packages, setup

install_requires = [
'torch>=1.3.0',
'gym>=0.9.7',
'numpy>=1.10.4',
'pillow',
'filelock',
"torch>=1.3.0",
"gym>=0.9.7",
"numpy>=1.10.4",
"pillow",
"filelock",
]

test_requires = [
'pytest',
'scipy',
'optuna',
'attrs<19.2.0', # pytest does not run with attrs==19.2.0 (https://github.com/pytest-dev/pytest/issues/3280) # NOQA
"pytest",
"scipy",
"optuna",
"attrs<19.2.0", # pytest does not run with attrs==19.2.0 (https://github.com/pytest-dev/pytest/issues/3280) # NOQA
]

setup(name='pfrl',
version='0.4.0',
description='PFRL, a deep reinforcement learning library',
long_description=codecs.open('README.md', 'r', encoding='utf-8').read(),
long_description_content_type='text/markdown',
author='Yasuhiro Fujita',
author_email='fujita@preferred.jp',
license='MIT License',
packages=find_packages(),
install_requires=install_requires,
test_requires=test_requires)
setup(
name="pfrl",
version="0.4.0",
description="PFRL, a deep reinforcement learning library",
long_description=codecs.open("README.md", "r", encoding="utf-8").read(),
long_description_content_type="text/markdown",
author="Yasuhiro Fujita",
author_email="fujita@preferred.jp",
license="MIT License",
packages=find_packages(),
install_requires=install_requires,
test_requires=test_requires,
)
1 change: 0 additions & 1 deletion tests/wrappers_tests/test_atari_wrappers.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
"""Currently this script tests `pfrl.wrappers.atari_wrappers.FrameStack`
only."""


from unittest import mock

import gym
Expand Down
46 changes: 29 additions & 17 deletions tools/plot_scores.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,41 +2,53 @@
import os

import matplotlib
matplotlib.use('Agg') # Needed to run without X-server

matplotlib.use("Agg") # Needed to run without X-server
import matplotlib.pyplot as plt
import pandas as pd


def main():
parser = argparse.ArgumentParser()
parser.add_argument('--title', type=str, default='')
parser.add_argument('--file', action='append', dest='files',
default=[], type=str,
help='specify paths of scores.txt')
parser.add_argument('--label', action='append', dest='labels',
default=[], type=str,
help='specify labels for scores.txt files')
parser.add_argument("--title", type=str, default="")
parser.add_argument(
"--file",
action="append",
dest="files",
default=[],
type=str,
help="specify paths of scores.txt",
)
parser.add_argument(
"--label",
action="append",
dest="labels",
default=[],
type=str,
help="specify labels for scores.txt files",
)
args = parser.parse_args()

assert len(args.files) > 0
assert len(args.labels) == len(args.files)

for fpath, label in zip(args.files, args.labels):
if os.path.isdir(fpath):
fpath = os.path.join(fpath, 'scores.txt')
fpath = os.path.join(fpath, "scores.txt")
assert os.path.exists(fpath)
scores = pd.read_csv(fpath, delimiter='\t')
plt.plot(scores['steps'], scores['mean'], label=label)
scores = pd.read_csv(fpath, delimiter="\t")
plt.plot(scores["steps"], scores["mean"], label=label)

plt.xlabel('steps')
plt.ylabel('score')
plt.legend(loc='best')
plt.xlabel("steps")
plt.ylabel("score")
plt.legend(loc="best")
if args.title:
plt.title(args.title)

fig_fname = args.files[0] + args.title + '.png'
fig_fname = args.files[0] + args.title + ".png"
plt.savefig(fig_fname)
print('Saved a figure as {}'.format(fig_fname))
print("Saved a figure as {}".format(fig_fname))


if __name__ == '__main__':
if __name__ == "__main__":
main()