Larger batches
This commit is contained in:
@@ -46,11 +46,11 @@ class LLBuffer(deque):
|
|||||||
|
|
||||||
def train(
|
def train(
|
||||||
expert_data,
|
expert_data,
|
||||||
expert_batch_size=3072,
|
expert_batch_size=4096,
|
||||||
discriminator_updates_per_round=20,
|
discriminator_updates_per_round=20,
|
||||||
generator_steps=1024,
|
generator_steps=1024,
|
||||||
generator_batch_size=1024,
|
generator_batch_size=1024,
|
||||||
generator_total_steps=4096,
|
generator_total_steps=8192,
|
||||||
generator_updates_per_round=10,
|
generator_updates_per_round=10,
|
||||||
discount=1.0,
|
discount=1.0,
|
||||||
epochs=200,
|
epochs=200,
|
||||||
@@ -138,7 +138,7 @@ def evaluate():
|
|||||||
|
|
||||||
# %%
|
# %%
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
with open("data/NormalizedIntersimpleExpertMu.001N100000_TLNRasterizedRouteRandomAgentLocationw70h70mppx1mu.001rskips50.pkl", "rb") as f:
|
with open("data/NormalizedIntersimpleExpertMu.001N50000_TLNRasterizedRouteRandomAgentLocationw70h70mppx1mu.001rskips50.pkl", "rb") as f:
|
||||||
trajectories = pickle.load(f)
|
trajectories = pickle.load(f)
|
||||||
transitions = rollout.flatten_trajectories(trajectories)
|
transitions = rollout.flatten_trajectories(trajectories)
|
||||||
train(transitions)
|
train(transitions)
|
||||||
|
|||||||
Reference in New Issue
Block a user