Full Width [alt+shift+f] Shortcuts [alt+shift+k]
Sign Up [alt+shift+s] Log In [alt+shift+l]
38

An entire Social Network in 1.6GB (GraphD Part 2)

from Jaz's Blog [alt+shift+b] in AI

In Part 1 of this series, we tried to answer the question “who do you follow who also follows user B” in Bluesky, a social network with millions of users and hundreds of millions of follow relationships. At the conclusion of the post, we’d developed an in-memory graph store for the network that uses HashMaps and HashSets to keep track of the followers of every user and the set of users they follow, allowing bidirectional lookups, intersections, unions, and other set operations for combining social graph data. I received some helpful feedback after that post where several people pointed me towards Roaring Bitmaps as a potential improvement on my implementation. They were right, Roaring Bitmaps would be an excellent fit for my Graph service, GraphD, and could also provide me with a much needed way to quickly persist and load the Graph data to and from disk on startup, hopefully reducing the startup time of the service. What are Bitmaps? If you just want to dive into the Roaring Bitmap spec, you can read the paper here, but it might be easier to first talk about bitmaps in general. You can think of a bitmap as a vector of one-bit values (like booleans) that let you encode a set of integer values. For instance, say we have 10,000 users on our website and want to keep track of which users have validated their email addresses. We could do this by creating a list of the uint32 user IDs of each user, in which case if all 10,000 users have validated their emails we’re storing 10k * 32 bits = 40KB. Or, we could create a vector of single-bit values that’s 10,000 bits long (10k / 8 = 1.25KB), then if a user has confirmed their email we can set the value at the index of their UID to 1. If we want to create a list of all the UIDs of validated accounts, we can walk the vector and record the index of each non-zero bit. If we want to check if user n has validated their email, we can do a O(1) lookup in the bitmap by loading the bit at index n and checking if it’s set. When Bitmaps...
20th Apr 2024

Stay updated

Get a weekly newsletter with the top 5 articles worth reading every week.

More from Jaz's Blog

Reverse Twins
2nd Oct 2025 • 24 votes
Turning Billions of Strings into Integers Every Second Without Collisions

I’ve recently started building a POC of a Redis RESP3 Wire Compatible Key/Value Database built on FoundationDB with @calabro.io and though it’s rather early, it’s already spawned a fun distributed systems problem that I thought would be interesting to share. Previously I’ve written about how I implemented a Graph DB via Roaring Bitmaps, representing relations as a bidirectional pair of sets. To support such use-cases in this new database, we’d like to represent sets of keys such that you can perform boolean operations on them (intersection, union, difference) relatively quickly even for very large sets (with millions of members). Supporting Larger Keys In the original Graph DB, we were representing user DID strings as uint32 UIDs to allow us to store millions of edge lists in very little space (e.g. the set of users who follow bsky.app) while being able to perform boolean operations between lists quickly (using Roaring Bitmaps’ parallel boolean operators). Since we were graphing follows, blocks, and other such User-to-User relationships, there was a practical maximum for the total number of user IDs in the low billions. We’ve continued exploring objects and relationships we’d like to represent as a Graph, and have realized that if we wanted to store e.g. the URIs of all posts a user has liked so we can intersect it with other users’ likes, we’re going to need a bigger keyspace! There are well over 15 Billion records in the AT Proto Ecosystem, each with a unique AT URI! Now our desired keyspace is much larger than can be represented by uint32 values and so we need to expand to uint64. Easy enough, let’s use the uint64 flavor of Roaring Bitmaps and simply intern URIs and User DIDs as uint64s, problem solved, right? Not quite… Interning Many Things at Once The AT Proto Firehose has hit historic peak traffic of over 1,500 evt/sec. We want to design a system that will handle many times more scale than we’ve ever seen in reality. This means designing for 10x or 100x would require us to be able to intern 15k to 150k new URIs per second into uint64 integers. Sounds easy enough, what’s the holdup? Well, in FoundationDB we’re able to use Transactions to do things like atomically increment a sequence safely when many other threads may be trying to do the same thing. This is simple enough to do in Go, we can just toss together a little helper function to acquire a new UID for our string: func (s *server) allocateNewUID(span trace.Span, tx fdb.Transaction) (uint64, error) { var newUID uint64 val, err := tx.Get(fdb.Key("last_uid")).Get() if err != nil { return 0, return fmt.Errorf("failed to get last UID: %w", err) } if len(val) == 0 { newUID = 1 // start from 1 } else { lastUID, err := strconv.ParseUint(string(val), 10, 64) if err != nil { return 0, return fmt.Errorf("failed to parse last UID: %w", err) } newUID = lastUID + 1 } tx.Set(fdb.Key("last_uid"), []byte(strconv.FormatUint(newUID, 10))) return newUID, nil } This function gets called from a fdb.Transaction which gets assigned a Transaction ID, then stages its changes, then tries to commit them. In FoundationDB, if your transaction is reading or modifying data written to by a different Transaction that finishes while you’re in-progress, your Transaction is thrown out and must be retried. For our UID assignment use-case, this is pretty problematic. We want to assign hundreds of thousands of new UIDs per second but if they’re all modifying the same key, concurrent transactions will constantly run into contention on the same data and will be forced to retry over and over again. This problem gets worse the more concurrent transactions you have trying to read from or write to the same key. Even if we stick to sequential access, if it takes ~5-10ms to assign a UID, we can only assign ~100-200 UIDs per second, nowhere near the throughput we need to support. How can we get past this problem and allow us to give strings unique uint64 UIDs in a high throughput and highly concurrent manner? Attempt #1: xxHash My first attempt to solve this problem was to try something that required no coordination and hash the string keys into uint64s using xxHash. xxHash is a non-cryptographic hash algorithm that supports incredibly high throughput (dozens of GB/sec) and can produce 64 bit unsigned integer hashes of strings trivially. Implementing this would look something like: Hash the incoming string key Lookup the uint64 UID to see if we’ve already assigned it to a string Reject the transaction if there’s a collision and give up Store the key in the UID map and the UID in the key map Use the UID for anything else we need While the uint64 keyspace is plenty large for our needs assuming we distribute evenly among the whole space, using a hashing algorithm with no coordination means there’s room for collisions and thus we’d need some additional logic (potentially by bucketing the keys somehow). Consulting the Birthday Problem we can see that a keyspace with 64 bit hashes has a >50% chance of containing a single collision when we have only ~5 billion keys in the set! That’s barely more keys than we can cram into a uint32 and definitely won’t suffice for the number of keys we expect to be storing! So, xxHash, while nice and coordination-free is probably not going to be the solution we need. What else can we do? Attempt #2: Billions of Sequences Incrementing one sequence is clearly not an option because we can only increment a single sequence ~100-200 times per second, but what if we instead had more than one sequence? Roaring Bitmaps managed to make highly efficient bitmap representations by breaking up a uint32 keyspace into a uint16-wide set of uint16-wide keyspaces. Can we do something similar here? Here’s an idea, what if we had just over 4 billion difference sequences and just picked one at random when we needed to assign a UID? Since we’re constructing our UIDs as a uint64, we can split the full UID into a pair of uint32s where the most-significant-bits are used to identify the sequence ID and the least-significant-bits are used to identify the value assigned to the UID within the sequence. So in our implementation, we get ~4.3 Billion sequence IDs that each have ~4.3 Billion incrementing values. As an example, if we were to randomly select Sequence ID 37 and then we increment that sequence to the value 5, we’d assemble the ID as 37<<32 + 5 which looks like 158,913,789,952 + 5 -> 158,913,789,957. Looking at the next Seuqence ID, we’d see 38 which, when left shifted by 32 gives us 163,208,757,248. You can see there’s a gap of ~4.3 billion values between the first UID assigned by each Sequence ID. Assuming we can increment a single sequence ~100 times per second with contention, we’re able to mint 430 Billion new UIDs per second without locking up (assuming the cluster can keep up). Storing ~4.3 billion sequences may be a bit expensive, but thankfully this strategy can scale up and down by picking a larger or smaller prefix size. If we only wanted to store say, ~16k sequences, we can pick a 14 bit prefix instead of a 32 bit prefix and then use a 50 bit sequence number. That spreads the load across 2^14 sequence IDs and significantly reduces storage requirements for Sequences. What does this look like in code? Well, it’s honestly not very complex! const uidSequencePrefix = "uid_sequence/" func (s *server) allocateNewUID(tx fdb.Transaction) (uint64, error) { // sequenceNum is the random uint32 sequence we are using for this allocation var sequenceNum uint32 var sequenceKey string // assignedUID is the uint32 within the sequence we will assign var assignedUID uint32 // Try up to 5 times to find a sequence that is not exhausted for range 5 { // Pick a random uint32 as the sequence we will be using for this UID sequenceNum = rand.Uint32() sequenceKey = fmt.Sprintf("%s%d", uidSequencePrefix, sequenceNum) val, err := tx.Get(fdb.Key(sequenceKey)).Get() if err != nil { return 0, fmt.Errorf("failed to get last UID: %w", err) } if len(val) == 0 { assignedUID = 1 // Start each sequence at 1 } else { lastUID, err := strconv.ParseUint(string(val), 10, 32) if err != nil { return 0, fmt.Errorf("failed to parse last UID: %w", err) } // If we have exhausted this sequence, pick a new random sequence if lastUID >= 0xFFFFFFFF { continue } assignedUID = uint32(lastUID) + 1 } } // If we failed to find a sequence after 5 tries, return an error if assignedUID == 0 { return 0, fmt.Errorf("failed to allocate new UID after 5 attempts") } // Assemble the 64-bit UID from the sequence ID and assigned UID newUID := (uint64(sequenceNum) << 32) | uint64(assignedUID) // Store the assigned UID back to the sequence key for the next allocation tx.Set(fdb.Key(sequenceKey), []byte(strconv.FormatUint(uint64(assignedUID), 10))) // Return the full 64-bit UID return newUID, nil } And there we go! We can now intern billions of strings per second with little to no contention in a distributed system while completely avoiding collisions and making full use of our keyspace! Conclusion Often times when designing distributed systems, patterns and strategies you see in seemingly unrelated libraries can inspire an elegant solution to the problem at hand. In the case of distributed, high-throughput string interning, horizontal scaling can be achieved by breaking up one large keyspace that requires strict coordination into billions of smaller keyspaces that can be randomly load-balanced across. Both patterns used in this technique are present elsewhere: Breaking up a large keyspace into a bunch of smaller keyspaces is present in Roaring Bitmaps (among other systems) Letting randomness and large numbers spread out resource contention is present in many load balancing systems This is one of my favorite parts of growing as an engineer: the more systems and strategies you familiarize yourself with, the more material you have to draw from when designing something new. Personal News A bit of personal news for y’all if you made it this far. Today is my last day as a member of the Bluesky team! The past 2+ years building out Bluesky’s Infrastructure and Platform team and scaling Bluesky from 100,000 -> 40,000,000 users have been the most intense and rewarding years of my life. I don’t have the words to express how much I’ve valued my time on the team and how much I care for the people I’ve worked with in what feels like a decade of real time. I’ve got some new adventures ahead and am excited to be embarking on a new journey within the next month (still building large-scale infrastructure, don’t worry). I plan to continue being involved in the AT Proto Community and to contribute to some cool projects other folks on the Bluesky team are working on from the FOSS space (like KVDB). To the team, I wish you all the best and will dearly miss getting to work with you all every day, but nothing lasts forever and I will always cherish the time I got to spend building an incredible platform with incredible people. If you’re interested in joining a world-class team doing important work, check out Bluesky’s open job listings here. There should be a new role opening up for a seasoned Go Engineer on the Platform team soon!

26th Sep 2025 • 1 votes
When Imperfect Systems are Good, Actually: Bluesky's Lossy Timelines

Often when designing systems, we aim for perfection in things like consistency of data, availability, latency, and more. The hardest part of system design is that it’s difficult (if not impossible) to design systems that have perfect consistency, perfect availability, incredibly low latency, and incredibly high throughput, all at the same time. Instead, when we approach system design, it’s best to treat each of these properties as points on different axes that we balance to find the “right fit” for the application we’re supporting. I recently made some major tradeoffs in the design of Bluesky’s Following Feed/Timeline to improve the performance of writes at the cost of consistency in a way that doesn’t negatively affect users but reduced P99s by over 96%. Timeline Fanout When you make a post on Bluesky, your post is indexed by our systems and persisted to a database where we can fetch it to hydrate and serve in API responses. Additionally, a reference to your post is “fanned out” to your followers so they can see it in their Timelines. This process involves looking up all of your followers, then inserting a new row into each of their Timeline tables in reverse chronological order with a reference to your post. When a user loads their Timeline, we fetch a page of post references and then hydrate the posts/actors concurrently to quickly build an API response and let them see the latest content from people they follow. The Timelines table is sharded by user. This means each user gets their own Timeline partition, randomly distributed among shards of our horizontally scalable database (ScyllaDB), replicated across multiple shards for high availability. Timelines are regularly trimmed when written to, keeping them near a target length and dropping older post references to conserve space. Hot Shards in Your Area Bluesky currently has around 32 Million Users and our Timelines database is broken into hundreds of shards. To support millions of partitions on such a small number of shards, each user’s Timeline partition is colocated with tens of thousands of other users’ Timelines. Under normal circumstances with all users behaving well, this doesn’t present a problem as the work of an individual Timeline is small enough that a shard can handle the work of tens of thousands of them without being heavily taxed. Unfortunately, with a large number of users, some of them will do abnormal things like… well… following hundreds of thousands of other users. Generally, this can be dealt with via policy and moderation to prevent abusive users from causing outsized load on systems, but these processes take time and can be imperfect. When a user follows hundreds of thousands of others, their Timeline becomes hyperactive with writes and trimming occurring at massively elevated rates. This load slows down the individual operations to the user’s Timeline, which is fine for the bad behaving user, but causes problems to the tens of thousands of other users sharing a shard with them. We typically call this situation a “Hot Shard”: where some resident of a shard has “hot” data that is being written to or read from at much higher rates than others. Since the data on the shard is only replicated a few times, we can’t effectively leverage the horizontal scale of our database to process all this additional work. Instead, the “Hot Shard” ends up spending so much time doing work for a single partition that operations to the colocated partitions slow down as well. Stacking Latencies Returning to our Fanout process, let’s consider the case of Fanout for a user followed by 2,000,000 other users. Under normal circumstances, writing to a single Timeline takes an average of ~600 microseconds. If we sequentially write to the Timelines of our user’s followers, we’ll be sitting around for 20 minutes at best to Fanout this post. If instead we concurrently Fanout to 1,000 Timelines at once, we can complete this Fanout job in ~1.2 seconds. That sounds great, except it oversimplifies an important property of systems: tail latencies. The average latency of a write is ~600 microseconds, but some writes take much less time and some take much more. In fact, the P99 latency of writes to the Timelines cluster can be as high as 15 milliseconds! What does this mean for our Fanout? Well, if we concurrently write to 1,000 Timelines at once, statistically we’ll see 10 writes as slow as or slower than 15 milliseconds. In the case of timelines, each “page” of followers is 10,000 users large and each “page” must be fanned out before we fetch the next page. This means that our slowest writes will hold up the fetching and Fanout of the next page. How does this affect our expected Fanout time? Each “page” will have ~100 writes as slow as or slower than the P99 latency. If we get unlucky, they could all stack up on a single routine and end up slowing down a single page of Fanout to 1.5 seconds. In the worst case, for our 2,000,000 Follower celebrity, their post Fanout could end up taking as long as 5 minutes! That’s not even considering P99.9 and P99.99 latencies which could end up being >1 second, which could leave us waiting tens of minutes for our Fanout job. Now imagine how bad this would be for a user with 20,000,000+ Followers! So, how do we fix the problem? By embracing imperfection, of course! Lossy Timelines Imagine a user who follows hundreds of thousands of others. Their Timeline is being written to hundreds of times a second, moving so fast it would be humanly impossible to keep up with the entirety of their Timeline even if it was their full-time job. For a given user, there’s a threshold beyond which it is unreasonable for them to be able to keep up with their Timeline. Beyond this point, they likely consume content through various other feeds and do not primarily use their Following Feed. Additionally, beyond this point, it is reasonable for us to not necessarily have a perfect chronology of everything posted by the many thousands of users they follow, but provide enough content that the Timeline always has something new. Note in this case I’m using the term “reasonable” to loosely convey that as a social media service, there must be a limit to the amount of work we are expected to do for a single user. What if we introduce a mechanism to reduce the correctness of a Timeline such that there is a limit to the amount of work a single Timeline can place on a DB shard. We can assert a reasonable limit for the number of follows a user should have to have a healthy and active Timeline, then increase the “lossiness” of their Timeline the further past that limit they go. A loss_factor can be defined as min(reasonable_limit/num_follows, 1) and can be used to probabilistically drop writes to a Timeline to prevent hot shards. Just before writing a page in Fanout, we can generate a random float between 0 and 1, then compare it to the loss_factor of each user in the page. If the user’s loss_factor is smaller than the generated float, we filter the user out of the page and don’t write to their Timeline. Now, users all have the same number of “follows worth” of Fanout. For example with a reasonable_limit of 2,000, a user who follows 4,000 others will have a loss_factor of 0.5 meaning half the writes to their Timeline will get dropped. For a user following 8,000 others, their loss factor of 0.25 will drop 75% of writes to their Timeline. Thus, each user has a effective ceiling on the amount of Fanout work done for their Timeline. By specifying the limits of reasonable user behavior and embracing imperfection for users who go beyond it, we can continue to provide service that meets the expectations of users without sacrificing scalability of the system. Aside on Caching We write to Timelines at a rate of more than one million times a second during the busy parts of the day. Looking up the number of follows of a given user before fanning out to them would require more than one million additional reads per second to our primary database cluster. This additional load would not be well received by our database and the additional cost wouldn’t be worth the payoff for faster Timeline Fanout. Instead, we implemented an approach that caches high-follow accounts in a Redis sorted set, then each instance of our Fanout service loads an updated version of the set into memory every 30 seconds. This allows us to perform lookups of follow counts for high-follow accounts millions of times per second per Fanount service instance. By caching values which don’t need to be perfect to function correctly in this case, we can once again embrace imperfection in the system to improve performance and scalability without compromising the function of the service. Results We implemented Lossy Timelines a few weeks ago on our production systems and saw a dramatic reduction in hot shards on the Timelines database clusters. In fact, there now appear to be no hot shards in the cluster at all, and the P99 of a page of Fanout work has been reduced by over 90%. Additionally, with the reduction in write P99s, the P99 duration for a full post Fanout has been reduced by over 96%. Jobs that used to take 5-10 minutes for large accounts now take <10 seconds. Knowing where it’s okay to be imperfect lets you trade consistency for other desirable aspects of your systems and scale ever higher. There are plenty of other places for improvement in our Timelines architecture, but this step was a big one towards improving throughput and scalability of Bluesky’s Timelines. If you’re interested in these sorts of problems and would like to help us build the core data services that power Bluesky, check out this job listing. If you’re interested in other open positions at Bluesky, you can find them here.

19th Feb 2025 • 63 votes
Emoji Griddle
30th Oct 2024 • 35 votes
Jetstream: Shrinking the AT Proto Firehose by >99%

Bluesky recently saw a massive spike in activity in response to Brazil’s ban of Twitter. As a result, the AT Proto event firehose provided by Bluesky’s Relay at bsky.network has increased in volume by a huge amount. The average event rate during this surge increased by ~1,300%. Before this new surge in activity, the firehose would produce around 24 GB/day of traffic. After the surge, this volume jumped to over 232 GB/day! Keeping up with the full, verified firehose quickly became less practical on cheap cloud infrastructure with metered bandwidth. To help reduce the burden of operating bots, feed generators, labelers, and other non-verifying AT Proto services, I built Jetstream as an alternative, lightweight, filterable JSON firehose for AT Proto. How the Firehose Works The AT Proto firehose is a mechanism used to keep verified, fully synced copies of the repos of all users. Since repos are represented as Merkle Search Trees, each firehose event contains an update to the user’s MST which includes all the changed blocks (nodes in the path from the root to the modified leaf). The root of this path is signed by the repo owner, and a consumer can keep their copy of the repo’s MST up-to-date by applying the diff in the event. For a more in-depth explanation of how Merkle Trees are constructed, check out this explainer. Practically, this means that for every small JSON record added to a repo, we also send along some number of MST blocks (which are content-addressed hashes and thus very information-dense) that are mostly useful for consumers attempting to keep a fully synced, verified copy of the repo. You can think of this as the difference between cloning a git repo v.s. just grabbing the latest version of the files without the .git folder. In this case, the firehose effectively streams the diffs for the repository with commits, signatures, and metadata, which is inherently heavier than a point-in-time checkout of the repo. Because firehose events with repo updates are signed by the repo owner, they allow a consumer to process events from any operator without having to trust the messenger. This is the “Authenticated” part of the Authenticated Transfer (AT) Protocol and is crucial to the correct functioning of the network. That being said, of the hundreds of consumers of Bluesky’s production Relay, >90% of them are building feeds, bots, and other tools that don’t keep full copies of the entire network and don’t verify MST operations at all. For these consumers, all they actually process is the JSON records created, updated, and deleted in each event. If consumers already trust the provider to do validation on their end, they could get by with a much more lightweight data stream. How Jetstream Works Jetstream is a streaming service that consumes an AT Proto com.atproto.sync.subscribeRepos stream and converts it into lightweight, friendly JSON. If you want to try it out yourself, you can connect to my public Jetstream instance and view all posts on Bluesky in realtime: $ websocat "wss://jetstream2.us-east.bsky.network/subscribe?wantedCollections=app.bsky.feed.post" Note: the above instance is operated by Bluesky PBC and is free to use, more instances are listed in the official repo Readme Jetstream converts the CBOR-encoded MST blocks produced by the AT Proto firehose and translates them into JSON objects that are easier to interface with using standard tooling available in programming languages. Since Repo MSTs only contain records in their leaf nodes, this means Jetstream can drop all of the blocks in an event except for those of the leaf nodes, typically leaving only one block per event. In reality, this means that Jetstream’s JSON firehose is nearly 1/10 the size of the full protocol firehose for the same events, but lacks the verifiability and signatures included in the protocol-level firehose. Jetstream events end up looking something like: { "did": "did:plc:eygmaihciaxprqvxpfvl6flk", "time_us": 1725911162329308, "type": "com", "commit": { "rev": "3l3qo2vutsw2b", "type": "c", "collection": "app.bsky.feed.like", "rkey": "3l3qo2vuowo2b", "record": { "$type": "app.bsky.feed.like", "createdAt": "2024-09-09T19:46:02.102Z", "subject": { "cid": "bafyreidc6sydkkbchcyg62v77wbhzvb2mvytlmsychqgwf2xojjtirmzj4", "uri": "at://did:plc:wa7b35aakoll7hugkrjtf3xf/app.bsky.feed.post/3l3pte3p2e325" } }, "cid": "bafyreidwaivazkwu67xztlmuobx35hs2lnfh3kolmgfmucldvhd3sgzcqi" } } Each event lets you know the DID of the repo it applies to, when it was seen by Jetstream (a time-based cursor), and up to one updated repo record as serialized JSON. Check out this 10 second CPU profile of Jetstream serving 200k evt/sec to a local consumer: By dropping the MST and verification overhead by consuming from relay we trust, we’ve reduced the size of a firehose of all events on the network from 232 GB/day to ~41GB/day, but we can do better. Jetstream and zstd I recently read a great engineering blog from Discord about their use of zstd to compress websocket traffic to/from their Gateway service and client applications. Since Jetstream emits marshalled JSON through the websocket for developer-friendliness, I figured it might be a neat idea to see if we could get further bandwidth reduction by employing zstd to compress events we send to consumers. zstd has two basic operating modes, “simple” mode and “streaming” mode. Streaming Compression At first glance, streaming mode seems like it’d be a great fit. We’ve got a websocket connection with a consumer and streaming mode allows the compression to get more efficient over the lifetime of the connection. I went and implemented a streaming compression version of Jetstream where a consumer can request compression when connecting and will get zstd compressed JSON sent as binary messages over the socket instead of plaintext. Unfortunately, this had a massive impact on Jetstream’s server-side CPU utilization. We were effectively compressing every message once per consumer as part of their streaming session. This was not a scalable approach to offering compression on Jetstream. Additionally, Jetstream stores a buffer of the past 24 hours (configurable) of events on disk in PebbleDB to allow consumers to replay events before getting transitioned into live-tailing mode. Jetstream stores serialized JSON in the DB, so playback is just shuffling the bytes into the websocket without having to round-trip the data into a Go struct. When we layer in streaming compression, playback becomes significantly more expensive because we have to compress outgoing events on-the-fly for a consumer that’s catching up. In real numbers, this increased CPU usage of Jetstream by 23% while lowering the throughput of playback from ~200k evt/sec to ~28k evt/sec for a single local consumer. When in streaming mode, we can’t leverage the bytes we compress for one consumer and reuse them for another consumer because zstd’s streaming context window may not be in sync between the two consumers. They haven’t received exactly the same data in the session so the clients on the other end don’t have their state machines in the same state. Since streaming mode’s primary advantage is giving us eventually better efficiency as the encoder learns about the data, what if we just taught the encoder about the data at the start and compress each message statelessly? Dictionary Mode zstd offers a mechanism for initializing an encoder/decoder with pre-optimized settings by providing a dictionary trained on a sample of the data you’ll be encoding/decoding. Using this dictionary, zstd essentially uses it’s smallest encoded representations for the most frequently seen patterns in the sample data. In our case, where we’re compressing serialized JSON with a common event shape and lots of common property names, training a dictionary on a large number of real events should allow us to represent the common elements among messages in the smallest number of bytes. For take two of Jetstream with zstd, let’s to use a single encoder for the whole service that utilizes a custom dictionary trained on 100,000 real events. We can use this encoder to compress every event as we see it, before persisting and emitting it to consumers. Now we end up with two copies of every event, one that’s just serialized JSON, and one that’s statelessly compressed to zstd using our dictionary. Any consumers that want compression can have a copy of the dictionary on their end to initialize a decoder, then when we broadcast the shared compressed event, all consumers can read it without any state or context issues. This requires the consumers and server to have a pre-shared dictionary, which is a major drawback of this implementation but good enough for our purposes. That leaves the problem of event playback for compression-enabled clients. An easy solution here is to just store the compressed events as well! Since we’re only sticking the JSON records into our PebbleDB, the actual size of the 24 hour playback window is <8GB with sstable compression. If we store a copy of the JSON serialized event and a copy of the zstd compressed event, this will, at most, double our storage requirements. Then during playback, if the consumer requests compression, we can just shuffle bytes out of the compressed version of the DB into their socket instead of having to move it through a zstd encoder. Savings Running with a custom dictionary, I was able to get the average Jetstream event down from 482 bytes to just 211 bytes (~0.44 compression ratio). Jetstream allows us to live tail all posts on Bluesky as they’re posted for as little as ~850 MB/day, and we could keep up with all events moving through the firehose during the Brazil Twitter Exodus weekend for 18GB/day (down from 232GB/day). With this scheme, Jetstream is required to compress each event only once before persisting it to disk and emitting it to connected consumers. The CPU impact of these changes is significant in proportion to Jetstream’s incredibly light load but it’s a flat cost we pay once no matter how many consumers we have. (CPU profile from a 30 second pprof sample with 12 consumers live-tailing Jetstream) Additionally, with Jetstream’s shared buffer broadcast architecture, we keep memory allocations incredibly low and the cost per consumer on CPU and RAM is trivial. In the allocation profile below, more than 80% of the allocations are used to consume the full protocol firehose. The total resident memory of Jetstream sits below 16MB, 25% of which is actually consumed by the new zstd dictionary. To bring it all home, here’s a screenshot from the dashboard of my public Jetstream instance serving 12 consumers all with various filters and compression settings, running on a $5/mo OVH VPS. At our new baseline firehose activity, a consumer of the protocol-level firehose would require downloading ~3.16TB/mo to keep up. A Jetstream consumer getting all created, updated, and deleted records without compression enabled would require downloading ~400GB/mo to keep up. A Jetstream consumer that only cares about posts and has zstd compression enabled can get by on as little as ~25.5GB/mo, <99% of the full weight firehose. Feel free to join the conversation about Jetstream and zstd on Bluesky.

24th Sep 2024 • 41 votes

More in AI

Prompting as self-portrait

Weekly curated resources for designers — thinkers and makers.

an hour ago • 1 votes
Coming soon: New York City’s hearing on AI risks

New York tries to take care of its own

9 hours ago • 1 votes
Dyson CameraJet

So when I saw Dyson had a $500 toothbrush, I was excited. Finally, advertising that targets me! I love brushing my teeth, and I have more money than I know how to spend. Not because I’m particularly rich, but because most stuff doesn’t really appeal to me. Like if I owned a helicopter it would just be a headache, because like imagine one day I get a call from the hangar saying the hangar is flooding and the water is rising and you need to move your helicopter. I’m thousands of miles away and need a helicopter pilot in the next 30 minutes, a new place to store it, was the maintenance even done will we even be able to take off on short notice and really I just am upset with myself because I made the poor decision to purchase a helicopter, and once I come back to reality I feel relieved that I don’t own a helicopter and this scenario will never happen to me. I do however, by means of my birthday, own a Dyson CameraJet (pictured above). It broke within 30 seconds of the first brushing. None of the LEDs turn on anymore. I spent an hour investigating, finally opening the user removable battery compartment to find the Spearmint Dyson Low-foaming mouth rinse had leaked inside. And by how the toothbrush is designed, it’s clear the entire electronics compartment was flooded with the stuff. Here’s the top comment on Reddit about this toothbrush. Apparently this is happening to everyone, “a potential for water seepage” they say. Dyson wants me to find the receipt and return it through some obtuse process that probably doesn’t work, dude it was a gift I just want my $500 toothbrush to work. They claim they worked on it for 6 years, but it’s clear their QA Process doesn’t include putting any liquid in the device. It clearly should, ideally for all devices but at least for spot checks on some. It’s sad to see this. At comma, we put every comma four in a highly stressful environment for 16 hours, a superset of the state it’s in driving, while testing all peripherals: the camera, IMU, GPS, screen, etc… We have gotten the failure rate super low by doing this, and for the few that do fail it’s usually after a while. There’s no excuse for a mature consumer electronics company to not design a procedure to fully test the functionality of each device before shipping. This shows some serious dysfunction at the company, and they should take this as a wake up call to fix their processes and issue a recall for the toothbrush. Dyson, if you see this post, e-mail me when I can drop by the Dyson store in ifc mall Hong Kong and swap it for a new one. I don’t want a stupid process, I want a real technical explanation of the issue and a working fancy toothbrush.

3 days ago • 1 votes
Why do OpenAI's GPT-2 weights beat mine? Part five: data quality

When I finished learning how to build an LLM from scratch, I was left with a mystery: my own models were not as good as OpenAI's original GPT-2 models, despite being based on the same architecture. My models all had 163M parameters, and followed the design from Sebastian Raschka's book "Build a Large Language Model (from Scratch)". That meant that they were pretty much the same as the setup for the OpenAI GPT-2 "small" instance, except that they did not use weight-tying or bias on the QKV matrices. Weight-tying means that you re-use the initial embedding matrix as the output head at the end, and using it means that GPT-2 small saved quite a few parameters -- it was 124M rather than 163M -- at, at least in my own experiments, a cost in quality; similarly, while I found that QKV bias made a tiny improvement in loss terms, I'd felt it was likely within the noise. But GPT-2 small consistently beat my models on an instruction fine-tuning (IFT) task -- also adapted from Raschka's book. That test fine-tunes the model on a subset of the Alpaca dataset, until validation loss starts rising, and then runs a test set through the resulting model. The responses to the test set questions are stored, and then I run all of the responses from all of the models under test past GPT 5.5 in one go to get an aggregate score; more details here. GPT-2 small always did better than any of my models on this. Additionally, it did surprisingly well on a simpler eval -- one that just measured the cross entropy loss it got on a test set. It scored close to my own best models, and better than many of them. What made this result particularly interesting was that the test set in question was a split of my own training data; my models would not have seen it when training (at least, in theory), but it seems likely that it would be much more similar to their own training data than it was to OpenAI's. I've checked two things while probing this mystery: It seems very likely that the GPT-2 models were overtrained by modern standards; would overtraining my own models get them closer? It turned out that no, it probably didn't help with the IFT eval (though there might have been some signal there). It did help quite a lot with the test loss eval, though. The way I was handling dropout in the IFT test might have been unduly benefiting some models while working against others. I decided to standardise on not using dropout during this eval, as (counter-intuitively for me) it seemed to harm the results of most models, even those that had been pre-trained with dropout. In particular, the OpenAI weights were harmed by using dropout, and making a change that benefited them (along with some of my own models) seemed the most conservative approach to take in investigating this. The next thing I wanted to look into was the training data. The exact dataset that the various GPT-2 models were trained on has never been released; all we know about it is from the paper, where they say: [W]e created a new web scrape which emphasizes document quality. To do this we only scraped web pages which have been curated/filtered by humans. Manually filtering a full web scrape would be exceptionally expensive so as a starting point, we scraped all outbound links from Reddit, a social media platform, which received at least 3 karma. This can be thought of as a heuristic indicator for whether other users found the link interesting, educational, or just funny. They called it "WebText". There is an OpenWebText that tries to replicate it, but although they tried to follow the same procedure as the original, there's no guarantee that it is all that similar. By comparison, I'd normally been training against FineWeb. While this is a general web-scraping dataset, without the "curation" provided by using only stuff that was linked from upvoted Reddit posts, it has been refined to remove any obvious junk. I had felt that it was pretty much equivalent. But what if I were wrong about that? I decided to see if I could get better models by using better data. The starting point Here's a table of all of the models I've been comparing to date. The "Test loss" column shows how well the model in question did on that held-back cross entropy loss evaluation. The "IFT epochs" column shows how many epochs of fine-tuning the model needed before its validation loss started rising, the "IFT score" the score that GPT 5.5 gave the model's responses to the test set of my Alpaca data, and the "IFT rank" the model's rank in terms of that score. The OpenAI small model is in there in bold, and I've also included the OpenAI medium model for comparison purposes. Test loss IFT epochs IFT score IFT rank OpenAI weights: medium 3.231442 2 43.75 1 JAX, overtrained one long epoch 3.324953 3 19.77 4 JAX, overtrained two normal epochs 3.326482 4 19.72 5 JAX, with MHA bias, no dropout 3.418784 4 18.69 6 JAX, no MHA bias, no dropout 3.420089 5 21.46 3 JAX, no MHA bias, with dropout 3.476802 5 13.22 15 OpenAI weights: small 3.499677 2 26.00 2 1xrtx3090-stacked-interventions 3.538161 4 13.77 14 8xa100m40-stacked-interventions-1 3.577761 4 10.76 18 Cloud FineWeb, 8x A100 40 GiB 3.673623 3 17.72 7 1xrtx3090-baseline 3.683835 4 15.74 8 8xa100m40-baseline 3.691526 3 14.19 13 Cloud FineWeb, 8x H100 80 GiB 3.724507 4 14.33 12 Cloud FineWeb, 8x A100 80 GiB 3.729900 3 11.34 17 Cloud FineWeb, 8x B200 160 GiB 3.771478 4 14.67 11 Local FineWeb train 3.943522 5 12.31 16 Local FineWeb-Edu extended train 4.134991 5 15.04 9 Local FineWeb-Edu train 4.166892 5 14.99 10 You can see that the OpenAI small model did pretty well in terms of the test loss, when you consider that it has 39M fewer weights than my models and was being tested against a dataset that differs more from its likely training data than it does from my own models'. Additionally, the specific models that did better than OpenAI's small one were all trained with JAX rather than PyTorch -- my hypothesis for that is that it's a result of the JAX ones getting better initial weights by pure chance. But the big difference was in the IFT score. In the specific run that gave the results in this table, the OpenAI small model got 26.00 -- the closest of my own models was more than 4.5 points lower, at 21.46. This difference was consistent over all of my other test runs. The GPT-2 small model was always ahead of mine. (GPT-2 medium, of course, beat GPT-2 small and all of my models, but given that it is twice the size of mine, that's not a big surprise.) Now, quite some time ago, I had tried looking into data quality as a lever to pull for model performance. At the bottom of the table, with the worst test loss of all models, you can see two models: "Local FineWeb-Edu train" "Local FineWeb-Edu extended train" These two were (as you might guess from the names) trained on the FineWeb-Edu dataset, which includes just the most "educational" data from FineWeb. They scored very badly on the test loss score. Given that the test dataset is from FineWeb, that's not a big surprise -- as I've written previously: If you train a model on Jane Austen and then evaluate against Chuck Tingle, then you're not going to get amazing results. But again, GPT-2 had the same issue, and did perfectly well on the test loss eval. On the other hand, while these FineWeb-Edu models' performance on the IFT eval wasn't stellar -- there are plenty of my other models ahead of them -- they did seem to punch above their weight. Consistently across all of the IFT evals I've done, they have scored higher than many of the others -- despite their poor loss on the test eval. Additionally: they were amongst the first models that I trained, before I'd spent time learning about how to optimise my hyperparameters and training loop. They did not use gradient clipping, they did use dropout, their batch size was just "whatever I could squeeze into the GPU", and I didn't set the learning rate to the right kind of value or schedule it over the course of the training run. So maybe a new training run on FineWeb-Edu plus my training improvements would help? And maybe some other tweaks to the training data would be worth looking into? The plan I decided to see what would happen if I trained some models with better-quality data. Specifically, I would train models with my current optimised loop and hyperparameters on four different datasets: FineWeb-Edu -- essentially the same as "Local FineWeb-Edu train" but with a better training setup. This would test the "more educational -> better" hypothesis. A 50:50 split of FineWeb and FineWeb-Edu. I've read that LLMs can be helped by having a decent amount of lower-quality data in their training loop, as it helps them to generalise. Perhaps having some FineWeb in there in addition to the FineWeb-Edu stuff would improve that test loss score while also helping the IFT test? A "curated" dataset containing 45% of its contents from FineWeb, 45% from FineWeb-Edu, and 10% from the Simple English Wikipedia. The full Wikipedia is huge, and full of obscure facts -- while the Simple English one is small and hopefully richer in useful information on a per-token basis. And conveniently, Answer.ai have made a snapshot of it available on Hugging Face Hub. Might deliberately putting a bunch of encyclopaedic data into the training set make the model better at the IFT eval (which has lots of factual questions in it, like "who wrote Pride and Prejudice")? OpenWebText. Even though I was unsure how well it matched the original WebText, given that it was there, it seemed silly to not try training something on it and see how it matched up. I would train each model on 3.2B tokens of the chosen dataset; that's the Chinchilla-optimal amount for my 163M-parameter models. If there were any interesting results, then I might consider doing overtrained models later on. I decided to be at least vaguely scientific about this, and to pre-register some predictions: The FineWeb-Edu-only model would do pretty badly on the test loss, but better than my older FineWeb-Edu models (90%). It would also punch above its weight on the IFT eval (90%). The 50:50 split: I expected it to do worse on the test eval than my JAX FineWeb-only models (70%), but better than the FineWeb-Edu one (90%). I wasn't sure about how it would do on the IFT eval, but thought it might be somewhere in between the two groups (60%). The curated dataset I had high hopes for in terms of the IFT eval -- let's say 80% chance of it being the best of all of my models. For the test loss eval, I expected it to do about as well as the 50:50 split, maybe a little bit worse (70%). I had no idea how the OpenWebText eval would do! Could be worse, could be better. Here's how things turned out. The FineWeb-Edu model I already had a dataset based on FineWeb-Edu ready to go, from when I trained those two original models. It is just the 10B-token sample of the original dataset at the time I generated it last December, formatted appropriately for my training script (details on the dataset card). I kicked off a training run with my JAX code (which I've been using for the other posts in this series): giles@poppy:~/Dev/jax-gpt2-from-scratch (main)$ XLA_PYTHON_CLIENT_MEM_FRACTION=0.95 uv run train.py full-llm-full-train-with-mha-output-bias-fineweb-edu datasets/ 2026-09-11 18:11:47.991583 Downloading dataset Fetching 4 files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 4/4 [00:00<00:00, 1772.93it/s] Download complete: : 0.00B [00:00, ?B/s] | 0/4 [00:00<?, ?it/s] 2026-09-11 18:11:48.226273 Loading dataset into RAM Download complete: : 0.00B [00:00, ?B/s] 2026-09-11 18:16:29.507646 Creating model 2026-09-11 18:16:33.042509 Creating optimizer 2026-09-11 18:16:34.138990 Start train 0%| | 0/33165 [00:00<?, ?it/s] 2026-09-11 18:17:38.486288 Saving checkpoint 1%|▌ | 173/33165 [13:22<39:17:03, 4.29s/it, loss=6.897, tps=21,201] ...and just less than 40 hours later, I had a model: Training complete in 142,912.226 seconds 2026-09-13 09:58:26.437276 Tokens seen: 3,260,252,160 2026-09-13 09:58:26.437284 Throughput: 22,813 tokens/second 2026-09-13 09:58:26.437302 Final train loss: 3.342 2026-09-13 09:58:26.437309 Done I converted the saved JAX safetensors file from the last checkpoint into a format that would be compatible with my PyTorch eval code, and ran my smoke test: how would it complete the sentence "Every effort moves you"? Every effort moves you closer to God’s Kingdom, and even closer to Him. As we can see in That was nice and coherent -- if unusually religious! -- so that was promising. I ran the test eval: giles@perry:~/Dev/ddp-base-model-from-scratch (main)$ uv run test_loss.py datasets/ ../jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-fineweb-edu/model.json ../jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-fineweb-edu/checkpoints/latest/pytorch-model.safetensors Fetching 4 files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 4/4 [00:00<00:00, 2758.50it/s] 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 3200/3200 [03:52<00:00, 13.74it/s] Loss against our test dataset: 3.632900 That was pretty good, putting it at a better test loss than all of the models I had trained without optimised hyperparameters, and worse than all of the ones I had trained on FineWeb with optimised hyperparameters. So that fit in with my prediction that it would be better than the old FineWeb-Edu models; the fact that it was also better than the non-optimised training runs with FineWeb seemed sensible enough that I felt silly for not having predicted that it would have fallen exactly there :-) I decided to leave the IFT eval until the end so that I could check all of the models from these experiments together, so it was time to upload this one to Hugging Face, and move on to the next model. 50:50 FineWeb to FineWeb-Edu I put together a new repo with a script to prepare datasets specifically for my training setup. You provide it with config that specifies some source datasets along with information about how to process them and how to mix them together, and it uploads a new dataset to Hugging Face Hub with the required characteristics. For example, for the 50:50 FineWeb to FineWeb-Edu split, the config looked like this: { "seed": 42, "tokens_desired": 10000000000, "upload_dataset_name": "gpjt/fw-fwedu-5050-gpt2-tokens", "sources": [ { "name": "FineWeb", "hf_id": "HuggingFaceFW/fineweb", "hf_name": "sample-10BT", "hf_split": "train", "item_field": "text", "weight": 50 }, { "name": "FineWeb-Edu", "hf_id": "HuggingFaceFW/fineweb-edu", "hf_name": "sample-10BT", "hf_split": "train", "item_field": "text", "weight": 50 } ] } The way the script works is pretty simple: it works out (based on those weights and the tokens_desired) how many tokens it wants from each source dataset, shuffles the items in the sources, then it loops until it has the desired number of tokens or more stored in an output. In the loop, it works out which source is currently most under-represented, grabs an item from it, tokenises it, and adds it to the output. Running it with that 50:50 config seemed to work fine: giles@perry:~/Dev/prepare-llm-training-dataset (main)$ uv run prepare-dataset.py runs/fw-fwedu-5050/ Resolving data files: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████| 27468/27468 [00:00<00:00, 89875.56it/s] Loading dataset shards: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████| 102/102 [00:00<00:00, 133.75it/s] Resolving data files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 2410/2410 [00:00<00:00, 87461.48it/s] Loading dataset shards: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████| 98/98 [00:00<00:00, 200.09it/s] 2026-09-13 20:13:22.000187: Generating dataset; per-source counts 2026-09-13 20:13:22.000217: FineWeb: 5,000,000,000 2026-09-13 20:13:22.000221: FineWeb-Edu: 5,000,000,000 FineWeb: 100%|████████████████████████████████████████████████████████████████████████████████████████████████▉| 4999999705/5000000000 [1:01:33<00:00, 1353639.33token/s] FineWeb-Edu: 5000000363token [1:01:33, 1353639.47token/s] 2026-09-13 21:14:55.747239: Done generating tokens 2026-09-13 21:14:55.748480: FineWeb: 4,999,999,705 / 5,000,000,000 (1.000, 1 iterators) 2026-09-13 21:14:55.748487: FineWeb-Edu: 5,000,000,363 / 5,000,000,000 (1.000, 1 iterators) 2026-09-13 21:14:55.748489: Total: 10,000,000,068 2026-09-13 21:14:55.748491: Catting... 2026-09-13 21:16:29.565152: Catted into a tensor of shape torch.Size([10000000068]) 2026-09-13 21:16:29.566663: Saving... 2026-09-13 21:16:36.006267: Saved 2026-09-13 21:16:36.009413: Uploading to gpjt/fw-fwedu-5050-gpt2-tokens Processing Files (1 / 1) : 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0GB / 20.0GB, 117MB/s New Data Upload : 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 14.6GB / 14.6GB, 98.1MB/s ...du-5050/train.safetensors: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0GB / 20.0GB 2026-09-13 21:17:59.545875: Done So we had almost-perfect 50:50 balance between the datasets, and it saved this dataset on Hugging Face. I ran a script to double-check that it looked sane, and it did, so it was time to spin up a training run: giles@perry:~/Dev/jax-gpt2-from-scratch (main)$ XLA_PYTHON_CLIENT_MEM_FRACTION=0.90 uv run train.py full-llm-full-train-with-mha-output-bias-fw-fwedu-5050 datasets/ 2026-09-13 21:20:59.880918 Downloading dataset Fetching 2 files: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 2/2 [01:13<00:00, 36.70s/it] Download complete: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0G/20.0G [01:13<00:00, 1.24GB/s] 2026-09-13 21:22:13.521745 Loading dataset into RAM Download complete: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0G/20.0G [01:13<00:00, 272MB/s] 2026-09-13 21:22:33.787720 Creating model 2026-09-13 21:22:35.501063 Creating optimizer 2026-09-13 21:22:36.043837 Start train 0%| | 0/33165 [00:00<?, ?it/s] 2026-09-13 21:23:11.437206 Saving checkpoint 0%| | 26/33165 [02:20<38:07:05, 4.14s/it, loss=9.308, tps=18,246] That was running on perry, my normal workstation, and I kicked it off in parallel with the "curated" model training run below on poppy my training box, but I'll keep the runs separate for the purposes of this writeup. When this had been running for an hour or so, our power went out. My guess is that having the tumble dryer running, the car charging, the kettle boiling, the electric hob switched on, and two machines doing training runs is a bit too much for our electrics... which might be a problem in the future, especially if (as planned) I make poppy a multi-GPU machine. However, as things stand, I was able to kick it off again after switching the circuit breaker back on, and things held up. Again, about 40 hours later: Training complete in 136,060.457 seconds 2026-09-15 12:05:26.432638 Tokens seen: 3,227,516,928 2026-09-15 12:05:26.432642 Throughput: 23,721 tokens/second 2026-09-15 12:05:26.432650 Final train loss: 3.793 2026-09-15 12:05:26.432653 Done (Note that the numbers reported at the end of a restarted run like this only include what happened after the restart.) I converted it to PyTorch-compatible tensors, and did the smoke test: Every effort moves you on to other options—in fact, it’s not even worth that effort. Just make Looking good! Time for the loss test: giles@perry:~/Dev/ddp-base-model-from-scratch (main)$ uv run test_loss.py datasets/ ../jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-fw-fwedu-5050/model.json ../jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-fw-fwedu-5050/checkpoints/latest/pytorch-model.safetensors Fetching 4 files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 4/4 [00:00<00:00, 1192.07it/s] 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 3200/3200 [03:53<00:00, 13.72it/s] Loss against our test dataset: 3.462454 That was almost in keeping with my prediction that it would do worse than the JAX FineWeb-only models, except that it was better than the worst of those, "JAX, no MHA bias, with dropout": it was actually better than I predicted. So, a promising model. Time to upload it to Hugging Face -- and now let's move on to the next one. The "curated" dataset With my dataset-preparation script, this was easy enough to set up: { "seed": 42, "tokens_desired": 10000000000, "upload_dataset_name": "gpjt/fw-fwedu-simplewiki-gpt2-tokens", "sources": [ { "name": "FineWeb", "hf_id": "HuggingFaceFW/fineweb", "hf_name": "sample-10BT", "hf_split": "train", "item_field": "text", "weight": 45 }, { "name": "FineWeb-Edu", "hf_id": "HuggingFaceFW/fineweb-edu", "hf_name": "sample-10BT", "hf_split": "train", "item_field": "text", "weight": 45 }, { "name": "Simple English Wikipedia", "hf_id": "answerdotai/simplewiki", "hf_name": "articles", "hf_split": "train", "item_field": "md", "weight": 10 } ] } Running that worked nicely: giles@perry:~/Dev/prepare-llm-training-dataset (main)$ uv run prepare-dataset.py runs/fw-fwedu-simplewiki/ Resolving data files: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████| 27468/27468 [00:00<00:00, 90196.13it/s] Loading dataset shards: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████| 102/102 [00:00<00:00, 358.90it/s] Resolving data files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 2410/2410 [00:00<00:00, 88254.11it/s] Loading dataset shards: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████| 98/98 [00:00<00:00, 589.23it/s] 2026-09-13 18:59:04.106327: Generating dataset; per-source counts 2026-09-13 18:59:04.106387: FineWeb: 4,500,000,000 2026-09-13 18:59:04.106407: FineWeb-Edu: 4,500,000,000 2026-09-13 18:59:04.106422: Simple English Wikipedia: 1,000,000,000 FineWeb: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████▉| 4499997964/4500000000 [59:41<00:00, 1256362.56token/s] FineWeb-Edu: 4500000607token [59:41, 1256363.31token/s] Simple English Wikipedia: 1000002889token [59:41, 279192.58token/s] 2026-09-13 19:58:45.874744: Done generating tokens 2026-09-13 19:58:45.876043: FineWeb: 4,499,997,964 / 4,500,000,000 (1.000, 1 iterators) 2026-09-13 19:58:45.876048: FineWeb-Edu: 4,500,000,607 / 4,500,000,000 (1.000, 1 iterators) 2026-09-13 19:58:45.876052: Simple English Wikipedia: 1,000,002,889 / 1,000,000,000 (1.000, 6 iterators) 2026-09-13 19:58:45.876054: Total: 10,000,001,460 2026-09-13 19:58:45.876056: Catting... 2026-09-13 20:00:18.811748: Catted into a tensor of shape torch.Size([10000001460]) 2026-09-13 20:00:18.813169: Saving... 2026-09-13 20:00:22.773873: Saved 2026-09-13 20:00:22.773936: Uploading to gpjt/fw-fwedu-simplewiki-gpt2-tokens Processing Files (1 / 1) : 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0GB / 20.0GB, 143MB/s New Data Upload : 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 19.8GB / 19.8GB, 142MB/s ...plewiki/train.safetensors: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0GB / 20.0GB 2026-09-13 20:01:59.270021: Done One thing that is worth noting in that output is the "6 iterators" for the Simple English Wikipedia. If a source dataset runs out of items while we're building up the results in this script, we start iterating over it again (with a different seed for the shuffle so that the ordering is different). The "6 iterators" means that it needed to do that 6 times -- the original creation of the iterator at the start of the script, and five more. So that means that the Simple English Wikipedia is repeated (oversampled) somewhere between five and six times in the dataset. That's not a bad thing! From what I've read, it's actually quite standard to oversample highly educational content in LLM training datasets. And anyway, the dataset the script generated was 10B tokens, of which we're only using 3.2B for the training run in this post, so it would only appear somewhere between one and two times. The repetition would likely only really cut in if and when we did an overtrained model on the dataset. Anyway, I ran my check against the uploaded dataset -- the first few items were clearly from FineWeb, FineWeb-Edu, and the Simple English Wikipedia. It was time to kick off a training run: giles@poppy:~/Dev/jax-gpt2-from-scratch (main)$ XLA_PYTHON_CLIENT_MEM_FRACTION=0.95 uv run train.py full-llm-full-train-with-mha-output-bias-fw-fwedu-simplewiki datasets/ 2026-09-13 20:24:48.037024 Downloading dataset Downloading (incomplete total...): 0.00B [00:00, ?B/s] Warning: You are sending unauthenticated requests to the HF Hub. Please set a HF_TOKEN to enable higher rate limits and faster downloads. | 0/2 [00:00<?, ?it/s] WARNING:huggingface_hub.utils._http:Warning: You are sending unauthenticated requests to the HF Hub. Please set a HF_TOKEN to enable higher rate limits and faster downloads. Fetching 2 files: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 2/2 [02:51<00:00, 85.85s/it] Download complete: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0G/20.0G [02:51<00:00, 435MB/s] 2026-09-13 20:27:39.934884 Loading dataset into RAM Download complete: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████████| 20.0G/20.0G [02:51<00:00, 116MB/s] 2026-09-13 20:31:20.492877 Creating model 2026-09-13 20:31:24.054143 Creating optimizer 2026-09-13 20:31:25.100832 Start train 0%| | 0/33165 [00:00<?, ?it/s] 2026-09-13 20:32:29.650379 Saving checkpoint 0%|▎ | 107/33165 [08:38<39:05:39, 4.26s/it, loss=7.631, tps=20,293] Again, this was interrupted by the power outage that hit the 50:50 training run, but I was able to restart from a checkpoint. After another 22 hours, it crashed with an error that I've seen before: jax.errors.JaxRuntimeError: INTERNAL: CUDA error: Failed to end stream capture: CUDA_ERROR_STREAM_CAPTURE_INVALIDATED: operation failed due to a previous error during capture [executable_name='jit_train_step'] I put it aside as a one-off oddity when I hit it last time, but this time I dug in a bit more. I noted that it had not ever happened on perry, but seemed to be an issue on poppy, and that poppy had an older version of CUDA and the Nvidia drivers -- might that be the cause? I decided to upgrade those before kicking off the next run, but for now just restarted the run from the most recent checkpoint. (Note for anyone who is hitting the same error: it has not occurred since the upgrade, so that's worth trying.) This time it completed OK: Training complete in 59,564.515 seconds 2026-09-15 15:56:52.909888 Tokens seen: 1,367,212,032 2026-09-15 15:56:52.909894 Throughput: 22,953 tokens/second 2026-09-15 15:56:52.909912 Final train loss: 3.332 2026-09-15 15:56:52.909959 Done Again, these numbers just show what happened after the most recent restart. I copied it over to perry, converted it into a format that was compatible with my PyTorch code, and ran the smoke test: Every effort moves you by the air, for it will make you a better athlete, so your body becomes bigger and stronger Coherent enough -- time for the loss eval: giles@perry:~/Dev/ddp-base-model-from-scratch (main)$ uv run test_loss.py datasets/ ~/Dev/jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-fw-fwedu-simplewiki/model.json ~/Dev/jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-fw-fwedu-simplewiki/checkpoints/latest/pytorch-model.safetensors Fetching 4 files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 4/4 [00:00<00:00, 1007.64it/s] 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 3200/3200 [03:57<00:00, 13.48it/s] Loss against our test dataset: 3.542460 Again, in line with my predictions -- worse than the JAX FineWeb-only models, and indeed than the very best PyTorch one, 1xrtx3090-stacked-interventions, and also worse than the 50:50 split, but better than the FineWeb-Edu one. I uploaded it to Hugging Face, and it was time to move on to what was meant to be the final model for this set of experiments. The OpenWebText run Again, this was a simple enough config to set up: { "seed": 42, "tokens_desired": 10000000000, "upload_dataset_name": "gpjt/openwebtext-gpt2-tokens", "sources": [ { "name": "OpenWebText", "hf_id": "Skylion007/openwebtext", "hf_name": "plain_text", "hf_split": "train", "item_field": "text", "weight": 50 } ] } ...and the build and upload process worked well (and took much less time -- for some reason, sampling randomly from a single dataset is faster than sampling from two or three): giles@perry:~/Dev/prepare-llm-training-dataset (main)$ uv run prepare-dataset.py runs/openwebtext/ Resolving data files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████| 80/80 [00:00<00:00, 32723.26it/s] Resolving data files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████| 80/80 [00:00<00:00, 97940.55it/s] Loading dataset shards: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████| 80/80 [00:00<00:00, 1200.13it/s] 2026-09-15 13:16:47.622617: Generating dataset; per-source counts 2026-09-15 13:16:47.622645: OpenWebText: 10,000,000,000 Resolving data files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████| 80/80 [00:00<00:00, 45602.65it/s] Resolving data files: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████| 80/80 [00:00<00:00, 67650.06it/s] Loading dataset shards: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████| 80/80 [00:00<00:00, 307.11it/s] OpenWebText: 10000000024token [31:46, 5246208.64token/s] 2026-09-15 13:48:33.761350: Done generating tokens 2026-09-15 13:48:33.762021: OpenWebText: 10,000,000,024 / 10,000,000,000 (1.000, 2 iterators) 2026-09-15 13:48:33.762026: Total: 10,000,000,024 2026-09-15 13:48:33.762028: Catting... 2026-09-15 13:49:33.115508: Catted into a tensor of shape torch.Size([10000000024]) 2026-09-15 13:49:33.115923: Saving... 2026-09-15 13:49:36.365978: Saved 2026-09-15 13:49:36.366027: Uploading to gpjt/openwebtext-gpt2-tokens Processing Files (0 / 1) : 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████▉| 20.0GB / 20.0GB, 147MB/s New Data Upload : 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████| 19.9GB / 19.9GB, 147MB/s ...webtext/train.safetensors: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████▉| 20.0GB / 20.0GB 2026-09-15 13:51:16.202890: Done Note that it needed to oversample -- that "2 iterators". OpenWebText is about 40 GiB uncompressed, and so that's about 10B GPT-2 tokens -- presumably just a little bit less. Again, given that I was planning to use just the first 3.2B tokens of the dataset, I didn't feel that it would matter. I ran the check script on the newly-uploaded Hugging Face dataset and all looked well, so that was all set for the training run. I upgraded poppy first with a sudo pacman -Syu to see if that helped with the weird error that I got in the previous run (which, as I said, it looks like it did), then kicked it off: giles@poppy:~/Dev/jax-gpt2-from-scratch (main)$ XLA_PYTHON_CLIENT_MEM_FRACTION=0.95 uv run train.py full-llm-full-train-with-mha-output-bias-openwebtext datasets/ 2026-09-15 16:42:32.606185 Downloading dataset Fetching 2 files: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 2/2 [00:00<00:00, 941.38it/s] Download complete: : 0.00B [00:00, ?B/s] | 0/2 [00:00<?, ?it/s] 2026-09-15 16:42:32.879987 Loading dataset into RAM Download complete: : 0.00B [00:00, ?B/s] 2026-09-15 16:45:40.438791 Creating model 2026-09-15 16:45:43.840269 Creating optimizer 2026-09-15 16:45:44.848351 Start train 0%| | 0/33165 [00:00<?, ?it/s] 2026-09-15 16:46:50.632075 Saving checkpoint 1%|█ | 332/33165 [24:33<38:45:54, 4.25s/it, loss=6.623, tps=22,154] About 31 hours in, it crashed again, but this time it was my own dumb fault: poppy has a relatively small disk and I ran out of space. I fixed that and kicked it off again from the most recent checkpoint, and this time it completed: Training complete in 33,927.995 seconds 2026-09-17 11:25:10.835989 Tokens seen: 779,747,328 2026-09-17 11:25:10.835994 Throughput: 22,982 tokens/second 2026-09-17 11:25:10.836012 Final train loss: 3.165 2026-09-17 11:25:10.836018 Done I converted it to PyTorch for the smoke test: Every effort moves you through each phase, so it's not a complete picture. I'm sure your story was ...which looked solid, so it was time for the test loss eval: giles@perry:~/Dev/ddp-base-model-from-scratch (main)$ uv run test_loss.py datasets/ ~/Dev/jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-openwebtext/model.json ~/Dev/jax-gpt2-from-scratch/runs/full-llm-full-train-with-mha-output-bias-openwebtext/checkpoints/latest/pytorch-model.safetensors Fetching 4 files: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 4/4 [00:00<00:00, 674.76it/s] 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 3200/3200 [03:59<00:00, 13.37it/s] Loss against our test dataset: 4.045255 Our worst score yet in this experiment! Worse than any of my models so far, apart from the two FineWeb-Edu ones I did without optimised hyperparameters. Now, the first draft of this post went straight to the results from here, but the story wasn't quite over yet... Test set contamination GPT-6 Astra is relentless. Before I publish any of these posts, I run them past an editorial board of LLMs to look for issues. GPT-6 Astra not only checked the text, it also visited the code I'd linked to to check that out too, and spotted something problematic. It's obvious in retrospect, but my code to build the new datasets had a high risk of including the contents of the -- in theory held-back -- test set. The way that the test set was generated was that I downloaded the 10B sample of FineWeb back in December, splitting it into 99% training data and 1% "validation". That validation split was about 100M tokens, and I was only using the first 19M or so for actual validation runs during training, so I (somewhat arbitrarily) designated about 19M other tokens starting at position 50M in there as my test set. Now, my new dataset-generation code was just sampling randomly from the complete 10B sample of FineWeb. So there was nothing stopping it from pulling in data that was in that old validation split! That meant that it was quite likely that my new "curated" and "50:50" datasets contained at least some of the test set that was meant to have been held back from the models during training. On reflection, the problem was potentially even worse. FineWeb-Edu is a subset of FineWeb; my existing FineWeb-Edu dataset came from the 10B sample of the Hugging Face original, and so it also could potentially contain documents that I'd put into the test set. The first thing to do was to establish the size of the problem. I wrote a script to take in a "forbidden" dataset and split; this was assumed to be formatted as one big tensor of GPT-2 tokens, which is what all of my datasets are. It would then split it by end-of-text tokens, and generate a hash and a token count for each resulting "document". Optionally, you could restrict it to only considering a subset -- the n tokens starting at position p -- and it would then generate hashes/lengths for the documents inside that slice, or that overlapped it at the start or the end. I ran that to generate a list of hashes for the entire validation set -- the validation split of gpjt/fineweb-gpt2-tokens -- and then used a second script to check my various training sets (and the validation set itself) to see how much of a contamination problem there was. I got these results: Dataset Split Contamination with validation set gpjt/fineweb-gpt2-tokens validation 102163003 out of 102163003 tokens (100.00%) gpjt/fineweb-gpt2-tokens train 636166 out of 102163003 tokens (0.62%) gpjt/fineweb-edu-gpt2-tokens train 672189 out of 102163003 tokens (0.66%) gpjt/fw-fwedu-5050-gpt2-tokens train 49224580 out of 102163003 tokens (48.18%) gpjt/fw-fwedu-simplewiki-gpt2-tokens train 44233824 out of 102163003 tokens (43.30%) gpjt/openwebtext-gpt2-tokens train 212 out of 102163003 tokens (0.00%) So: The validation set was 100% "contaminated" with itself, which was a useful sanity check. The training set of gpjt/fineweb-gpt2-tokens had what I felt was a small level of contamination. It was interesting that there was any at all -- I think that must mean that there are some repeated documents in the original dataset, and some of them wound up with copies in both my training and validation splits. The gpjt/fineweb-edu-gpt2-tokens dataset also had what felt like a reassuringly low level of contamination. Both gpjt/fw-fwedu-5050-gpt2-tokens and gpjt/fw-fwedu-simplewiki-gpt2-tokens, however, looked problematic. In both cases, the training datasets had more than 40% of the validation/test set in them. gpjt/openwebtext-gpt2-tokens was, as you'd expect, almost completely uncontaminated. It looks like maybe one document happened to have been picked up by both the OpenWebText and the FineWeb crawls and then included in the bit of FineWeb I was using for validation. However, these numbers -- while scary, at least for the 50:50 and the curated datasets -- were not quite the ones to use. They showed how much of the full validation set showed up in the full training set; what I actually cared about was how much of the test set -- those 19M tokens starting at position 50M in the validation split -- was in the actual subset of the training datasets that I actually trained on -- the first ~3.2B of them. I re-ran the script to generate hashes for just the test set, and then re-ran the contamination-checking script, telling it just to look at the appropriate subset of the training tokens, and got this: Dataset (first 3.2B tokens only) Split Contamination with test set gpjt/fineweb-gpt2-tokens train 26557 out of 19632681 tokens (0.14%) gpjt/fineweb-edu-gpt2-tokens train 32079 out of 19632681 tokens (0.16%) gpjt/fw-fwedu-5050-gpt2-tokens train 2986889 out of 19632681 tokens (15.21%) gpjt/fw-fwedu-simplewiki-gpt2-tokens train 2682430 out of 19632681 tokens (13.66%) gpjt/openwebtext-gpt2-tokens train None It was clear that there was a problem -- certainly with gpjt/fw-fwedu-5050-gpt2-tokens and gpjt/fw-fwedu-simplewiki-gpt2-tokens. They'd seen what felt like a significant amount of the test set while training, so their results on the test loss eval were dubious at best. I decided to train those two models afresh, and see what the result was in terms of loss. If the difference was huge, I'd look into the risks of the (much smaller) contamination of gpjt/fineweb-gpt2-tokens and gpjt/fineweb-edu-gpt2-tokens. But if it was pretty small, I'd not worry about that too much. I extended the script that prepared datasets so that the config file could specify a forbidden_dataset. Any documents in the source datasets that matched forbidden ones would be excluded from the output. I then updated the config for gpjt/fw-fwedu-5050-gpt2-tokens and gpjt/fw-fwedu-simplewiki-gpt2-tokens so that the whole validation split of gpjt/fineweb-gpt2-tokens was forbidden, and re-generated them. You can see the updated datasets here and here. Running the contamination-checker script against them showed that they were clear. I then re-did the full training runs for those models; the uncontaminated version of the 50:50 split model is here, and the curated one is here. And the good news: both of them actually did very slightly better at the test loss eval than their equivalents that had been trained on the contaminated data: Model Contaminated Test loss JAX, FineWeb/FineWeb-Edu 50:50 No 3.449257 JAX, FineWeb/FineWeb-Edu 50:50 Yes 3.462454 JAX, curated No 3.534068 JAX, curated Yes 3.542460 There are a number of possibilities that come to mind; perhaps learning from the test set just doesn't happen with tiny 163M models like this, or perhaps while the contaminated models were learning, the benefit they got from that was outweighed by the data that they got instead of the test set data being in some way better for training purposes, at least in terms of the loss eval. But anyway, I felt that if the effect of seeing more than 10% of the test set data during training was so tiny, then the effect of seeing less than 0.2% -- which is what the FineWeb-Edu model in this set of training runs had, as did all of my other FineWeb-only models from previous experiments -- would be even smaller and I'd disregard it. That was excellent news! I didn't need to start all of my experiments from scratch. For the rest of this post, I will include the numbers and results for the contaminated models as well as the uncontaminated ones -- they're interesting for several reasons -- but for future posts I'll skip the contaminated ones. So -- finally! -- let's start digging into the final results. Results Firstly, I think it's worth taking a look at all of the test loss results in context. Here they are in a table, with the new models in bold: Test loss OpenAI weights: medium 3.231442 JAX, overtrained one long epoch 3.324953 JAX, overtrained two normal epochs 3.326482 JAX, with MHA bias, no dropout 3.418784 JAX, no MHA bias, no dropout 3.420089 JAX, FineWeb/FineWeb-Edu 50:50 (uncontaminated) 3.449257 JAX, FineWeb/FineWeb-Edu 50:50 (contaminated) 3.462454 JAX, no MHA bias, with dropout 3.476802 OpenAI weights: small 3.499677 JAX, curated (uncontaminated) 3.534068 1xrtx3090-stacked-interventions 3.538161 JAX, curated (contaminated) 3.542460 8xa100m40-stacked-interventions-1 3.577761 JAX, FineWeb-Edu 3.632900 Cloud FineWeb, 8x A100 40 GiB 3.673623 1xrtx3090-baseline 3.683835 8xa100m40-baseline 3.691526 Cloud FineWeb, 8x H100 80 GiB 3.724507 Cloud FineWeb, 8x A100 80 GiB 3.729900 Cloud FineWeb, 8x B200 160 GiB 3.771478 Local FineWeb train 3.943522 JAX, openwebtext 4.045255 Local FineWeb-Edu extended train 4.134991 Local FineWeb-Edu train 4.166892 I think there's something very clear here: with the new models, the more FineWeb that was in the training mix, the better the model did on this eval. I think I might have been subconsciously expecting that in the predictions I did before running these experiments, but in retrospect it's so incredibly obvious that I feel silly for not mentioning it explicitly! But that tells us something interesting. From the description in the paper, whatever OpenAI did the GPT-2 training run on, it was not like FineWeb. It was probably more similar to OpenWebText -- and yet, that model was the one that performed the worst on this test eval, so if it is more like OpenWebText, there must be some other factor involved. But moving on for now: how about the IFT test -- the one that kicked off all of this work in the first place? I generated a set of IFT responses for all of the new models, and then ran them (plus responses for all of the other models on that table above) past GPT 5.5, and found that one of my new models was getting quite close to the original GPT-2 small weights! So I did four more runs, so that I could get an average. Here are the results -- the "IFT score" is the average across all five runs of the judge, and the "IFT rank" is based on that. The "IFT epochs" was from the original result-generation script. Test loss IFT epochs IFT score IFT rank OpenAI weights: medium 3.231442 2 42.36 1 JAX, overtrained one long epoch 3.324953 3 18.67 7 JAX, overtrained two normal epochs 3.326482 4 18.71 6 JAX, with MHA bias, no dropout 3.418784 4 17.90 8 JAX, no MHA bias, no dropout 3.420089 5 20.50 4 JAX, FineWeb/FineWeb-Edu 50:50 (uncontaminated) 3.449257 4 17.69 9 JAX, FineWeb/FineWeb-Edu 50:50 (contaminated) 3.462454 4 19.30 5 JAX, no MHA bias, with dropout 3.476802 5 13.02 21 OpenAI weights: small 3.499677 2 25.19 2 JAX, curated (uncontaminated) 3.534068 4 16.63 10 1xrtx3090-stacked-interventions 3.538161 4 13.51 19 JAX, curated (contaminated) 3.542460 4 13.58 18 8xa100m40-stacked-interventions-1 3.577761 4 10.19 24 JAX, FineWeb-Edu 3.632900 4 24.56 3 Cloud FineWeb, 8x A100 40 GiB 3.673623 3 16.59 11 1xrtx3090-baseline 3.683835 4 15.15 12 8xa100m40-baseline 3.691526 3 13.64 16 Cloud FineWeb, 8x H100 80 GiB 3.724507 4 13.59 17 Cloud FineWeb, 8x A100 80 GiB 3.729900 3 10.79 23 Cloud FineWeb, 8x B200 160 GiB 3.771478 4 13.70 15 Local FineWeb train 3.943522 5 11.87 22 JAX, openwebtext 4.045255 4 13.28 20 Local FineWeb-Edu extended train 4.134991 5 14.29 14 Local FineWeb-Edu train 4.166892 5 14.69 13 If you want to see the full numbers, they're below. The number that initially surprised me, and made me decide to do multiple LLM-judge runs was the one for the "JAX, FineWeb-Edu" model. In my first run it came in at 24.35 vs the OpenAI small weights' 24.93 -- so close that I wondered if it might even beat them on a re-run. However, in the further four runs its score was consistently lower than the OpenAI model's, and the gap extended a bit in some. So, was FineWeb-Edu the clear winner here? Perhaps. If you look at the contaminated/uncontaminated pairs, something interesting pops out. For the 50:50 mix, the model trained with the contaminated dataset got 19.30, and the one trained on the uncontaminated one got 17.69 -- a difference of 1.61. For the "curated" dataset, the situation was even more interesting: uncontaminated got 16.63, while contaminated got 13.58, a delta of 3.05 points. Remember, the contamination issue is about whether or not the model saw the held-back test set during training. It was an issue for the test loss that is based on that test set, but is entirely orthogonal to the IFT test. From the IFT perspective, both contaminated and uncontaminated models in each case saw training data that was -- in theory, at least -- essentially the same in terms of quality. Indeed, the uncontaminated run saw almost the same data in the same order as the contaminated one, except that some items were omitted, and then extra ones were added to the end. The purpose of this set of experiments was to see how data quality affected the results on the IFT test set. But in the case of the curated model, something that should be unrelated to data quality changed the results by 3.05 points! If something as simple as changing which data of the same quality the model is trained with can affect the IFT score so drastically, it makes it a bit harder to be certain as to whether or not data quality really had the effect we were looking for. On the other hand, the FineWeb-Edu model came in at 24.56, which is 4.06 points better than the 20.50 that the closest other model got -- more than the 3.05 points we see in difference between the two curated dataset models. And it's worth noting that the model with 20.50 is "JAX, no MHA bias, no dropout", which has a subtly different architecture -- no bias on the output projection of the multi-head attention blocks. A better comparison might be "JAX, with MHA bias, no dropout", which got a score of 17.90, for a whacking great difference of 6.66 points. I think that without doing a very large number of training runs on different datasets with different mixes, each one created with a different seed, it would be hard to work out exactly what is in the noise here and what is not. However, that would cost a lot in terms of time. I think that the best thing here is to chalk this up as a fairly decent indication that FineWeb-Edu improves matters for the IFT eval, but far from a certainty. But it's certainly worth noting that whatever the noise is, it has a range of at least 3.05 points -- and the FineWeb-Edu model is just 0.63 points short of GPT-2 small! So there could well be something there. Of course, we don't know whether that model got (by chance) the best possible balance of FineWeb-Edu tokens, and could never win -- or whether it got a bad balance and would actually beat GPT-2 with a better one. So that's certainly worth keeping in mind. As an aside, the result for the curated dataset really surprised me. I had expected that it would be the best one, simply because it almost certainly contained more facts. I took a look at its answers to the questions -- one possibility that came to mind might be that it would get better responses to questions like "What is the chemical symbol for chlorine" or "Who wrote Pride and Prejudice" than the others, but would fail on less knowledge-based tasks. But it was terrible at fact-based questions too: Name the author of 'Pride and Prejudice'. What is the periodic symbol for chlorine? As I understand it, many real-world training runs do include (often oversampled) amounts of highly educational training data like this model's dataset did. But perhaps the models that I'm training are just too small to be able to make use of the data they gained that way -- maybe doing things this way and expecting good results is like asking six-year-old children to memorise stuff before they've learned enough to be able to make use of it 1. It's worth noting that the GPT-2 small model also failed on those factual questions. Well, anyway: I think we have some useful results here, so let's work out what that means for next steps. Conclusion The results we got in these experiments point in two interesting directions. The perfect connection between the amount of FineWeb in the training set and the result on the (FineWeb-based) test loss eval, while perfectly obvious in retrospect, really does highlight how mysterious it is that the OpenAI small weights do so well on that test. The fact that FineWeb-Edu did well on the IFT test tells us that there does seem to be value in using richer training data -- though the less-spectacular results of the 50:50 mix and the curated one weaken that a bit, as does the indicator of what the noise due to data selection from equivalently high-quality datasets might be. The OpenWebText result I think I'll ignore, given that -- while in theory it should be similar to what OpenAI trained on -- there are no guarantees, and it might differ in non-obvious ways for non-obvious reasons. I think that the right direction to take this going forward is to separate these two angles. I should chase a higher IFT score, and then once I have nailed that down, I should see what (if anything) might allow me to get the resulting model to improve its test score. But I will need to make sure that whatever dataset I use, I use various "mixes" of it -- versions created with different random seeds. In my earlier experiments with overtraining, I did find that it didn't seem to improve the IFT results -- but it did improve the test loss. So perhaps identifying the right combination of other factors to boost the IFT score, then overtraining the result, might help? Of course, my overtraining tests were with FineWeb, so the connection might not hold up as well if the starting model (as seems likely) was trained on a different dataset. Also, while working through the results here, I've come to the conclusion that the set of models I'm using is a bit confusing -- there are now different hyperparameter settings, small architectural differences (the MHA bias thing), dropout settings during the pre-training, and now datasets. I think that's OK for now; I should see this part of this series as more ideation than actually running the proper experiments. But at the end, when I have some solid hypotheses with a reasonable amount of backup, I should start from scratch: a baseline model, then staged interventions to build up to what (hopefully) will be a model as good as GPT-2 small. Anyway, I'll wrap this one up here. I think that the next lever to pull is (perhaps surprisingly) going to be weight tying. I had previously kind of disregarded that as a possibility, but while I was working on this post, something popped into my mind. The OpenAI models were originally trained with weight tying. My codebase does actually support doing it -- but because I got the OpenAI weights I'm using from the code in "Build a Large Language Model (from Scratch)", when I'm running the IFT test, the weights are not actually tied! We load up a model that has separate but identical embedding and output head matrices, and then we fine-tune that. So those two matrices can vary independently during fine-tuning -- to put it another way, while GPT-2 small was pre-trained with 124M parameters, the IFT test is being done on a 163M-parameter version. Does that give them some non-obvious advantage? And would adding weight-tying to my own models help, either with or without the output heads being independent at fine-tuning time? Stay tuned :-) Appendix: all IFT judge runs Here are the numbers for all of the IFT judge runs, included for completeness. You can see that the LLM judge ranks models very consistently between runs, but there is variation -- that is, on some runs it's in what I think of as a "better mood" than others, and if that's the case, it will give better scores -- but it will give them almost consistently between models, so all of the models do better. Note that (unlike the table above) this one is sorted by the average IFT score rather than the test loss. Model Run 1 Run 2 Run 3 Run 4 Run 5 Average OpenAI weights: medium 42.24 42.16 42.95 41.83 42.61 42.36 OpenAI weights: small 24.93 24.96 25.39 25.01 25.66 25.19 JAX, FineWeb-Edu 24.35 24.55 24.3 24.68 24.9 24.56 JAX, no MHA bias, no dropout 20.5 19.9 20.76 21.25 20.07 20.50 JAX, FineWeb/FineWeb-Edu 50:50 (contaminated) 19.16 18.86 19.61 19.17 19.7 19.30 JAX, overtrained two normal epochs 18.47 18.29 19.17 18.69 18.91 18.71 JAX, overtrained one long epoch 18.04 18.71 19.62 18.41 18.57 18.67 JAX, with MHA bias, no dropout 17.49 17.35 18.33 17.73 18.62 17.90 JAX, FineWeb/FineWeb-Edu 50:50 (uncontaminated) 17.37 17.73 17.53 18.01 17.83 17.69 JAX, curated (uncontaminated) 16.77 16.03 17.3 16.08 16.96 16.63 Cloud FineWeb, 8x A100 40 GiB 16.44 16.23 17.14 16.62 16.54 16.59 1xrtx3090-baseline 14.85 15.07 15.19 15.14 15.51 15.15 Local FineWeb-Edu train 14.37 14.23 15.08 14.79 15 14.69 Local FineWeb-Edu extended train 14.4 14.07 13.82 14.56 14.61 14.29 Cloud FineWeb, 8x B200 160 GiB 13.37 13.05 13.85 13.67 14.57 13.70 8xa100m40-baseline 13.64 13.36 13.9 13.32 13.97 13.64 Cloud FineWeb, 8x H100 80 GiB 13.45 13.32 13.6 13.51 14.07 13.59 JAX, curated (contaminated) 13.09 13.48 13.95 13.23 14.15 13.58 1xrtx3090-stacked-interventions 13.37 13.11 14.04 13.84 13.17 13.51 JAX, openwebtext 12.88 12.7 13.74 13.53 13.53 13.28 JAX, no MHA bias, with dropout 13.19 12.86 12.98 12.85 13.24 13.02 Local FineWeb train 11.75 11.75 12.21 11.46 12.19 11.87 Cloud FineWeb, 8x A100 80 GiB 10.68 10.2 11.03 10.55 11.49 10.79 8xa100m40-stacked-interventions-1 9.44 9.79 10.84 10.2 10.66 10.19 A small boy asleep on his right side, the right arm stuck out, the right hand hanging limp over the edge of the bed. Through a round grating in the side of a box a voice speaks softly. "The Nile is the longest river in Africa and the second in length of all the rivers of the globe. Although falling short of the length of the Mississippi-Missouri, the Nile is at the head of all rivers as regards the length of its basin, which extends through 35 degrees of latitude …" At breakfast the next morning, "Tommy," some one says, "do you know which is the longest river in Africa?" A shaking of the head. "But don't you remember something that begins: The Nile is the …" "The - Nile - is - the - longest - river - in - Africa - and - the - second - in - length - of - all - the - rivers - of - the - globe …" The words come rushing out. "Although - falling - short - of …" "Well now, which is the longest river in Africa?" The eyes are blank. "I don't know." "But the Nile, Tommy." "The - Nile - is - the - longest - river - in - Africa - and - second …" "Then which river is the longest, Tommy?" Tommy burst into tears. "I don't know," he howls. Brave New World, Aldous Huxley ↩

4 days ago • 1 votes
📚 BoredReading

You seem to be enjoying this.

Join free to unlock everything.

Create free account

Already have an account? Sign in