Skip to content
This repository was archived by the owner on Apr 23, 2025. It is now read-only.
Merged
Show file tree
Hide file tree
Changes from 3 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
29 changes: 19 additions & 10 deletions Examples/LeNet-MNIST/main.swift
Original file line number Diff line number Diff line change
Expand Up @@ -31,24 +31,33 @@ let dataset = MNIST(batchSize: batchSize, on: device)

// The LeNet-5 model, equivalent to `LeNet` in `ImageClassificationModels`.
var classifier = Sequential {
Conv2D<Float>(filterShape: (5, 5, 1, 6), padding: .same, activation: relu)
AvgPool2D<Float>(poolSize: (2, 2), strides: (2, 2))
Conv2D<Float>(filterShape: (5, 5, 6, 16), activation: relu)
AvgPool2D<Float>(poolSize: (2, 2), strides: (2, 2))
Flatten<Float>()
Dense<Float>(inputSize: 400, outputSize: 120, activation: relu)
Dense<Float>(inputSize: 120, outputSize: 84, activation: relu)
Dense<Float>(inputSize: 84, outputSize: 10)
Conv2D<Float>(filterShape: (5, 5, 1, 6), padding: .same, activation: relu)
AvgPool2D<Float>(poolSize: (2, 2), strides: (2, 2))
Conv2D<Float>(filterShape: (5, 5, 6, 16), activation: relu)
AvgPool2D<Float>(poolSize: (2, 2), strides: (2, 2))
Flatten<Float>()
Dense<Float>(inputSize: 400, outputSize: 120, activation: relu)
Dense<Float>(inputSize: 120, outputSize: 84, activation: relu)
Dense<Float>(inputSize: 84, outputSize: 10)
}

var optimizer = SGD(for: classifier, learningRate: 0.1)

let trainingProgress = TrainingProgress()
var trainingLoop = TrainingLoop(
training: dataset.training,
validation: dataset.validation,
optimizer: optimizer,
lossFunction: softmaxCrossEntropy,
callbacks: [trainingProgress.update])
metrics: [.accuracy],
callbacks: [CSVLogger().log])

// Compute statistics only when last batch ends.
(trainingLoop.statisticsRecorder!).shouldCompute = {
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
(
_ batchIndex: Int, _ batchCount: Int, _ epochIndex: Int, _ epochCount: Int,
_ event: TrainingLoopEvent
) -> Bool in
return event == .batchEnd && batchIndex + 1 == batchCount
}

try! trainingLoop.fit(&classifier, epochs: epochCount, on: device)
3 changes: 1 addition & 2 deletions Examples/MobileNetV1-Imagenette/main.swift
Original file line number Diff line number Diff line change
Expand Up @@ -29,12 +29,11 @@ let dataset = Imagenette(batchSize: 64, inputSize: .resized320, outputSize: 224,
var model = MobileNetV1(classCount: 10)
let optimizer = SGD(for: model, learningRate: 0.02, momentum: 0.9)

let trainingProgress = TrainingProgress()
var trainingLoop = TrainingLoop(
training: dataset.training,
validation: dataset.validation,
optimizer: optimizer,
lossFunction: softmaxCrossEntropy,
callbacks: [trainingProgress.update])
metrics: [.accuracy])

try! trainingLoop.fit(&model, epochs: 10, on: device)
3 changes: 1 addition & 2 deletions Examples/MobileNetV2-Imagenette/main.swift
Original file line number Diff line number Diff line change
Expand Up @@ -29,12 +29,11 @@ let dataset = Imagenette(batchSize: 64, inputSize: .resized320, outputSize: 224,
var model = MobileNetV2(classCount: 10)
let optimizer = SGD(for: model, learningRate: 0.002, momentum: 0.9)

let trainingProgress = TrainingProgress()
var trainingLoop = TrainingLoop(
training: dataset.training,
validation: dataset.validation,
optimizer: optimizer,
lossFunction: softmaxCrossEntropy,
callbacks: [trainingProgress.update])
metrics: [.accuracy])

try! trainingLoop.fit(&model, epochs: 10, on: device)
3 changes: 1 addition & 2 deletions Examples/ResNet-CIFAR10/main.swift
Original file line number Diff line number Diff line change
Expand Up @@ -29,12 +29,11 @@ let dataset = CIFAR10(batchSize: 10, on: device)
var model = ResNet(classCount: 10, depth: .resNet56, downsamplingInFirstStage: false)
var optimizer = SGD(for: model, learningRate: 0.001)

let trainingProgress = TrainingProgress()
var trainingLoop = TrainingLoop(
training: dataset.training,
validation: dataset.validation,
optimizer: optimizer,
lossFunction: softmaxCrossEntropy,
callbacks: [trainingProgress.update])
metrics: [.accuracy])

try! trainingLoop.fit(&model, epochs: 10, on: device)
4 changes: 2 additions & 2 deletions Examples/VGG-Imagewoof/main.swift
Original file line number Diff line number Diff line change
Expand Up @@ -39,12 +39,12 @@ public func scheduleLearningRate<L: TrainingLoopProtocol>(
}
}

let trainingProgress = TrainingProgress()
var trainingLoop = TrainingLoop(
training: dataset.training,
validation: dataset.validation,
optimizer: optimizer,
lossFunction: softmaxCrossEntropy,
callbacks: [trainingProgress.update, scheduleLearningRate])
metrics: [.accuracy],
callbacks: [scheduleLearningRate])

try! trainingLoop.fit(&model, epochs: 90, on: device)
1 change: 1 addition & 0 deletions Support/FileSystem.swift
Original file line number Diff line number Diff line change
Expand Up @@ -39,4 +39,5 @@ public protocol File {
func read(position: Int, count: Int) throws -> Data
func write(_ value: Data) throws
func write(_ value: Data, position: Int) throws
func append(_ value: Data) throws
}
7 changes: 7 additions & 0 deletions Support/FoundationFileSystem.swift
Original file line number Diff line number Diff line change
Expand Up @@ -58,4 +58,11 @@ public struct FoundationFile: File {
// TODO: Incorporate file offset.
try value.write(to: location)
}

public func append(_ value: Data) throws {
Comment thread
xihui-wu marked this conversation as resolved.
let fileHandler = try FileHandle(forUpdating: location)
try fileHandler.seekToEnd()
try fileHandler.write(contentsOf: value)
try fileHandler.close()
}
}
6 changes: 4 additions & 2 deletions TrainingLoop/CMakeLists.txt
Original file line number Diff line number Diff line change
@@ -1,8 +1,10 @@
add_library(TrainingLoop
LossFunctions.swift
Metrics.swift
TrainingLoop.swift
TrainingProgress.swift
TrainingStatistics.swift)
Callbacks/StatisticsRecorder.swift
Callbacks/ProgressPrinter.swift
Callbacks/CSVLogger.swift)
target_link_libraries(TrainingLoop PUBLIC
ModelSupport)
set_target_properties(TrainingLoop PROPERTIES
Expand Down
70 changes: 70 additions & 0 deletions TrainingLoop/Callbacks/CSVLogger.swift
Original file line number Diff line number Diff line change
@@ -0,0 +1,70 @@
import Foundation
import ModelSupport

/// A callback-based handler for logging the statistics to CSV file.

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

nice doc comment. English nit:

Suggested change
/// A callback-based handler for logging the statistics to CSV file.
/// A callback-based handler for logging statistics to a CSV file.

I would consider not leading with “callback-based.” In fact, "logging" and "CSV File" are kind of implied by the name. So the best description would explain what's being logged. “Statistics” is good, but what kind fo statistics? Training statistics, maybe? Maybe this should be thought of as “a log file” rather than “a logger?”

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.

Updated the doc comment. Any reason why this is LogFile? (Logger makes sense to me since its a logger not a file)

@dabrahams dabrahams Sep 23, 2020

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

This thing can only be constructed to log to a path in the filesystem, and when we invoke its one public method, log, it appends to that file. Properly used, if you have multiple instances, each instance logs to a different file. So instances have a 1-1 correspondence to log files. The word Logger doesn't imply any of those things. A more general Logger might be an interesting abstraction, but it would have a different API.

If I see this code, I know exactly what's happening.

// No argument label needed, arguably, because when you construct a thing called “File” with
// a string, the string is obviously a path.
LogFile eventLog(s)

eventLog.append(blah, blah, blah)

With Logger and log, it's less clear.

public class CSVLogger {
public var path: String
Comment thread
xihui-wu marked this conversation as resolved.

let foundationFS: FoundationFileSystem
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
let foundationFile: FoundationFile
Comment thread
xihui-wu marked this conversation as resolved.
Outdated

/// Create an instance that log statistics during the training loop.
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
public init(withPath path: String = "run/log.csv") {
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
self.path = path

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

It seems very unlikely to me that we actually need to store path in addition to foundationFile. Consider whether it can/should be dropped because you can get it from foundationFile.

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.

Same. Removed foundationFile and kept the path.

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Hmm, I'm not sure you want to open and close the file for every line logged, though. (I am presuming that the foundationFile object keeps the file open)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

How did you “resolve” this comment?

self.foundationFS = FoundationFileSystem()
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
self.foundationFile = FoundationFile(path: path)
}

/// The callback used to hook into the TrainingLoop for logging statistics.
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
///
/// - Parameters:
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
/// - loop: The TrainingLoop where an event has occurred.
/// - event: The training or validation event that this callback is responding to.
public func log<L: TrainingLoopProtocol>(_ loop: inout L, event: TrainingLoopEvent) throws {
Comment thread
xihui-wu marked this conversation as resolved.
switch event {
case .batchEnd:
guard let epochIndex = loop.epochIndex, let epochCount = loop.epochCount,
let batchIndex = loop.batchIndex, let batchCount = loop.batchCount
else {
return

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Why would any of these things be nil, and why are we bailing out when they're nil? That should be explained in a comment. Ditto for stats below. Also, why are these two separate guard statements?

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.

Merged stats into the same guard statements.
These are designed as optionals in existing TrainingLoop. I think that's because it's not the must-have values to complete a training process. I added an inline comment.

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

OK, first, I would rather see something like

// These properties will be `nil` unless stats logging was requested

because it describes the situation at a semantic level rather than at the level of what some code did.

(also writing "No-Op" adds nothing to what is already very obvious from the code)

But that said, it seems very unlikely that the comment I'd like to see is true of any but the last property. All the others refer to values that have nothing to do with logging. So I want to know what causes epochIndex to be nil, for example.

It's a big design flaw in trainingLoop that it has so many optionals, and that makes this task more difficult, but I believe that's not the code you're working on(?)

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.

These variables were originally designed to be optionals to store temporary data in protocol: https://github.com/tensorflow/swift-models/blob/master/TrainingLoop/TrainingLoop.swift#L76

The generic TrainingLoop that implements the protocol does set all these optionals. My guess on why it was designed so is that it allows other TrainingLoops not setting them.

So the point on the comment is NOT "if stats logging doesn't request it then these properties will be nil", but "if these properties are nil No-op on the CSVLogger".

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Then you should delete the comment. The code very clearly says that all by itself, so the comment explains nothing.

I don't know if you're missing the point I'm trying to make, or you just disagree with it, but I'm doing this code review 100% for your benefit as a programmer. If the review process is blocking your progress, please feel free to just commit the changes, and decide separately about whether you want the feedback I'm giving you here. If you do, we can continue to discuss it.

}

guard let stats = loop.lastStatsLog else {
return
}

if !FileManager.default.fileExists(atPath: path) {
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
try foundationFS.createDirectoryIfMissing(at: String(path[..<path.lastIndex(of: "/")!]))
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
try writeHeader(stats: stats)
}
try writeDataRow(
Comment thread
xihui-wu marked this conversation as resolved.
epoch: "\(epochIndex + 1)/\(epochCount)",
batch: "\(batchIndex + 1)/\(batchCount)",
stats: stats)
default:
return
}
}

func writeHeader(stats: [(String, Float)]) throws {
let head: String = (["epoch", "batch"] + stats.map { $0.0 }).joined(separator: ", ")
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
do {
try head.write(toFile: path, atomically: true, encoding: .utf8)
} catch {
print("Unexpected error in writing header line: \(error).")
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
throw error
}
}

func writeDataRow(epoch: String, batch: String, stats: [(String, Float)]) throws {
let dataRow: Data = (
Comment thread
dabrahams marked this conversation as resolved.
Outdated
"\n" + ([epoch, batch] + stats.map { String($0.1) }).joined(separator: ", ")
Comment thread
xihui-wu marked this conversation as resolved.
Outdated
).data(using: .utf8)!
do {
try foundationFile.append(dataRow)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Isn't it important to do this atomically if there are multiple writers to the same log file?

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.

Does FileHandle support atomically write ? (Also, we won't expect multiple files written at same time I think)

} catch {
print("Unexpected error in writing data row: \(error).")
throw error
}
}
}
93 changes: 93 additions & 0 deletions TrainingLoop/Callbacks/ProgressPrinter.swift
Original file line number Diff line number Diff line change
@@ -0,0 +1,93 @@
// Copyright 2020 The TensorFlow Authors. All Rights Reserved.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.

import Foundation

let progressBarLength = 30

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

length in what unit? I suppose it's probably characters. Being a top-level declaration this should have a doc comment, and that would be a perfect place to put the answer. Is this the number of = signs, or the whole length printed, or…?


/// A callback-based handler for printing the training or validation progress.
public class ProgressPrinter {

/// Create an instance that prints progress during the training loop.
/// The progress contains a dynamic progress bar followed by statistics of metrics.
public init() {
}

/// The callback used to hook into the TrainingLoop for printing progress.
///
/// An example of the progress would be:
/// Epoch 1/12
/// 468/468 [==============================] - loss: 0.4819 - accuracy: 0.8513
/// 79/79 [==============================] - loss: 0.1520 - accuracy: 0.9521
///
/// - Parameters:
/// - loop: The TrainingLoop where an event has occurred.
/// - event: The training or validation event that this callback is responding to.
public func print<L: TrainingLoopProtocol>(_ loop: inout L, event: TrainingLoopEvent) throws {
switch event {
case .epochStart:
guard let epochIndex = loop.epochIndex, let epochCount = loop.epochCount else {
return
}

Swift.print("Epoch \(epochIndex + 1)/\(epochCount)")
case .batchEnd:
guard let batchIndex = loop.batchIndex, let batchCount = loop.batchCount else {
return
}

let progressBar = formatProgressBar(
progress: Float(batchIndex + 1) / Float(batchCount), length: progressBarLength)
var stats: String = ""
if let lastStatsLog = loop.lastStatsLog {
stats = formatStats(lastStatsLog)
}

Swift.print(
"\r\(batchIndex + 1)/\(batchCount) \(progressBar)\(stats)",
terminator: ""
)
fflush(stdout)
case .epochEnd:
Swift.print("")
case .validationStart:
Swift.print("")
default:
return
}
}

func formatProgressBar(progress: Float, length: Int) -> String {
let progressSteps = Int(round(Float(length) * progress))
let leading = String(repeating: "=", count: progressSteps)
let separator: String
let trailing: String
if progressSteps < progressBarLength {
separator = ">"
trailing = String(repeating: ".", count: progressBarLength - progressSteps - 1)
} else {
separator = ""
trailing = ""
}
return "[\(leading)\(separator)\(trailing)]"
}

func formatStats(_ stats: [(String, Float)]) -> String {
var result = ""
for stat in stats {
result += " - \(stat.0): \(String(format: "%.4f", stat.1))"
}
return result
}
}
Loading